已合并
feat: part2 完善HostCPU自定义算子运行时与常量折叠支持 #4538
duhua创建于 8月24日
feat: part2 完善HostCPU自定义算子运行时与常量折叠支持 #4538
已合并
duhua创建于 8月24日
共 37 个文件变更+2232-120
@@ -71,6 +71,9 @@ REGISTER_PROF_TYPE(AicpuHostCompute);
71REGISTER_PROF_TYPE(LaunchMixKernelWithHandle);71REGISTER_PROF_TYPE(LaunchMixKernelWithHandle);
72REGISTER_PROF_TYPE(LaunchMixKernelWithFlag);72REGISTER_PROF_TYPE(LaunchMixKernelWithFlag);
73REGISTER_PROF_TYPE(ExecuteCustomOp);73REGISTER_PROF_TYPE(ExecuteCustomOp);
74+REGISTER_PROF_TYPE(ExecuteCustomOpWithInferShape);
75+REGISTER_PROF_TYPE(ExecuteHostCustomOp);
76+REGISTER_PROF_TYPE(ExecuteHostCustomOpWithInferShape);
74REGISTER_PROF_NON_LAUNCH_TYPE(AICoreUpdateContext);77REGISTER_PROF_NON_LAUNCH_TYPE(AICoreUpdateContext);
75REGISTER_PROF_NON_LAUNCH_TYPE(AICpuUpdateContext);78REGISTER_PROF_NON_LAUNCH_TYPE(AICpuUpdateContext);
76REGISTER_PROF_NON_LAUNCH_TYPE(StaAutoUpdateContext);79REGISTER_PROF_NON_LAUNCH_TYPE(StaAutoUpdateContext);
@@ -476,6 +476,7 @@ target_link_libraries(ge_compiler
476 error_manager476 error_manager
477 unified_dlog477 unified_dlog
478 runtime_headers478 runtime_headers
479+ gert
479 aihac_symbolizer480 aihac_symbolizer
480 lowering481 lowering
481 -Wl,--as-needed482 -Wl,--as-needed
@@ -17,7 +17,10 @@
17#include "custom_op_factory.h"17#include "custom_op_factory.h"
18#include "common/ge_common/ge_types.h"18#include "common/ge_common/ge_types.h"
19#include "common/checker.h"19#include "common/checker.h"
20+#include "graph/ascend_string.h"
20#include "graph/custom_op/cast.h"21#include "graph/custom_op/cast.h"
22+#include "graph/debug/ge_attr_define.h"
23+#include "graph/utils/attr_utils.h"
21#include "lowering/kernel_run_context_builder.h"24#include "lowering/kernel_run_context_builder.h"
22#include "common/compile_profiling/ge_trace_wrapper.h"25#include "common/compile_profiling/ge_trace_wrapper.h"
23#include "common/thread_pool/thread_pool.h"26#include "common/thread_pool/thread_pool.h"
@@ -142,12 +145,26 @@ ge::Status CompileCustomOpSerially(const std::vector<CompileTask *> *tasks) {
142 return ge::SUCCESS;145 return ge::SUCCESS;
143}146}
144 147 
148+bool IsCustomOpExecOnHostCpu(const ge::OpDescPtr &op_desc) {
149+ if (!ge::CustomOpFactory::IsExistOp(ge::AscendString(op_desc->GetTypePtr()), ge::OpBackend::kHostCPU)) {
150+ return false;
151+ }
152+ std::string lowering_func;
153+ return ge::AttrUtils::GetStr(op_desc, ge::kAttrLowingFunc, lowering_func) &&
154+ (lowering_func == ge::kHostCpuCustomOpLowerFunc);
155+}
156+ 
145ge::Status AppendCompileTaskIfNeeded(const ge::NodePtr &node, std::vector<CompileTask> &compile_tasks) {157ge::Status AppendCompileTaskIfNeeded(const ge::NodePtr &node, std::vector<CompileTask> &compile_tasks) {
146 const auto op_type = node->GetType();158 const auto op_type = node->GetType();
147 const ge::AscendString op_type_ascend(op_type.c_str());159 const ge::AscendString op_type_ascend(op_type.c_str());
148 if (!ge::CustomOpFactory::IsExistOp(op_type_ascend, ge::OpBackend::kDevice)) {160 if (!ge::CustomOpFactory::IsExistOp(op_type_ascend, ge::OpBackend::kDevice)) {
149 return ge::SUCCESS;161 return ge::SUCCESS;
150 }162 }
163+ if (IsCustomOpExecOnHostCpu(node->GetOpDesc())) {
164+ GELOGD("skip compile for custom op execute on host cpu, op_name:%s, op_type:%s", node->GetName().c_str(),
165+ node->GetType().c_str());
166+ return ge::SUCCESS;
167+ }
151 GELOGI("during optimize whole graph, %s is custom op", op_type_ascend.GetString());168 GELOGI("during optimize whole graph, %s is custom op", op_type_ascend.GetString());
152 auto *const base_custom_op_ptr = ge::CustomOpFactory::CreateOrGetCustomOp(op_type_ascend, ge::OpBackend::kDevice);169 auto *const base_custom_op_ptr = ge::CustomOpFactory::CreateOrGetCustomOp(op_type_ascend, ge::OpBackend::kDevice);
153 if (base_custom_op_ptr == nullptr) {170 if (base_custom_op_ptr == nullptr) {
@@ -68,8 +68,7 @@ void CustomOpsKernelInfoStore::GetAllOpsKernelInfo(std::map<std::string, OpInfo>
68bool CustomOpsKernelInfoStore::CheckSupported(const OpDescPtr &op_desc, std::string &reason) const {68bool CustomOpsKernelInfoStore::CheckSupported(const OpDescPtr &op_desc, std::string &reason) const {
69 (void)reason;69 (void)reason;
70 GE_ASSERT_NOTNULL(op_desc);70 GE_ASSERT_NOTNULL(op_desc);
71- std::lock_guard<std::mutex> lock(mu_);71+ return CustomOpFactory::IsExistOp(AscendString(op_desc->GetTypePtr()), OpBackend::kDevice);
72- return op_info_map_.count(op_desc->GetType()) > 0;
73}72}
74} // namespace custom73} // namespace custom
75} // namespace ge74} // namespace ge
@@ -18,7 +18,11 @@
18#include "base/err_msg.h"18#include "base/err_msg.h"
19#include "framework/common/debug/ge_log.h"19#include "framework/common/debug/ge_log.h"
20#include "analyzer/analyzer.h"20#include "analyzer/analyzer.h"
21+#include "common/ge_common/ge_types.h"
21#include "graph/ge_context.h"22#include "graph/ge_context.h"
23+#include "graph/custom_op_factory.h"
24+#include "graph/ascend_string.h"
25+#include "graph/utils/attr_utils.h"
22#include "graph/utils/graph_utils.h"26#include "graph/utils/graph_utils.h"
23#include "graph/utils/node_utils.h"27#include "graph/utils/node_utils.h"
24#include "graph/utils/op_type_utils.h"28#include "graph/utils/op_type_utils.h"
@@ -68,6 +72,21 @@ bool ExecOnHostCpu(const OpDescPtr &op_desc) {
68 (type == ge::PARTITIONEDCALL);72 (type == ge::PARTITIONEDCALL);
69 return (!is_not_cpu_op);73 return (!is_not_cpu_op);
70}74}
75+ 
76+bool IsHostCpuCustomOp(const OpDescPtr &op_desc) {
77+ return (op_desc != nullptr) && CustomOpFactory::IsExistOp(AscendString(op_desc->GetTypePtr()), OpBackend::kHostCPU);
78+}
79+ 
80+void SetHostCpuCustomOp(const OpDescPtr &op_desc, OpInfo &matched_op_info) {
81+ op_desc->SetOpEngineName(kEngineNameCustom);
82+ op_desc->SetOpKernelLibName(kCustomOpKernelLibName);
83+ matched_op_info.engine = kEngineNameCustom;
84+ matched_op_info.opKernelLib = kCustomOpKernelLibName;
85+ (void)AttrUtils::SetStr(op_desc, kAttrLowingFunc, kHostCpuCustomOpLowerFunc);
86+ GELOGD("DNNEngineManager:Set kernel_lib %s, atomic engine %s, to node %s, LowerFunc %s",
87+ kCustomOpKernelLibName.c_str(), kEngineNameCustom.c_str(), op_desc->GetName().c_str(),
88+ kHostCpuCustomOpLowerFunc.c_str());
89+}
71} // namespace90} // namespace
72 91 
73DNNEngineManager::DNNEngineManager() : init_flag_(false) {}92DNNEngineManager::DNNEngineManager() : init_flag_(false) {}
@@ -336,11 +355,17 @@ std::string DNNEngineManager::GetDNNEngineName(const ge::NodePtr &node_ptr,
336 return "";355 return "";
337 }356 }
338 GE_IF_BOOL_EXEC(ExecOnHostCpu(op_desc), return GetHostCpuEngineName(op_infos, op_desc, matched_op_info));357 GE_IF_BOOL_EXEC(ExecOnHostCpu(op_desc), return GetHostCpuEngineName(op_infos, op_desc, matched_op_info));
358+ const bool is_host_cpu_custom_op = IsHostCpuCustomOp(op_desc);
339 std::map<std::string, std::string> unsupported_reasons;359 std::map<std::string, std::string> unsupported_reasons;
340 for (const auto &it : op_infos) {360 for (const auto &it : op_infos) {
341 if ((exclude_engines.find(it.engine) != exclude_engines.end()) && (!is_op_specified_engine)) {361 if ((exclude_engines.find(it.engine) != exclude_engines.end()) && (!is_op_specified_engine)) {
342 continue;362 continue;
343 }363 }
364+ if ((it.engine == kHostCpuEngineName) && is_host_cpu_custom_op) {
365+ matched_op_info = it;
366+ SetHostCpuCustomOp(op_desc, matched_op_info);
367+ return kEngineNameCustom;
368+ }
344 const auto &kernel_name = it.opKernelLib;369 const auto &kernel_name = it.opKernelLib;
345 auto kernel_info_store = OpsKernelManager::GetInstance().GetOpsKernelInfoStore(kernel_name);370 auto kernel_info_store = OpsKernelManager::GetInstance().GetOpsKernelInfoStore(kernel_name);
346 if (kernel_info_store == nullptr) {371 if (kernel_info_store == nullptr) {
@@ -399,6 +424,18 @@ std::string DNNEngineManager::GetDNNEngineName(const ge::NodePtr &node_ptr,
399 }424 }
400 }425 }
401 426 
427+ // Fallback for host cpu custom op: when the op is registered as host cpu custom op but kHostCpuEngineName
428+ // is not in op_infos, and all other engines failed CheckSupported, we should still use the custom host
429+ // engine as a fallback.
430+ if (is_host_cpu_custom_op) {
431+ const bool host_cpu_excluded =
432+ (exclude_engines.find(kHostCpuEngineName) != exclude_engines.end()) && (!is_op_specified_engine);
433+ if (!host_cpu_excluded) {
434+ SetHostCpuCustomOp(op_desc, matched_op_info);
435+ return kEngineNameCustom;
436+ }
437+ }
438+ 
402 // concat unsupported reasons analyzed data selection439 // concat unsupported reasons analyzed data selection
403 std::string reason;440 std::string reason;
404 for (const auto &it : unsupported_reasons) {441 for (const auto &it : unsupported_reasons) {
@@ -610,6 +647,10 @@ std::string DNNEngineManager::GetCompositeEngineKernelLibName(const std::string
610 647 
611std::string DNNEngineManager::GetHostCpuEngineName(const std::vector<OpInfo> &op_infos, const OpDescPtr &op_desc,648std::string DNNEngineManager::GetHostCpuEngineName(const std::vector<OpInfo> &op_infos, const OpDescPtr &op_desc,
612 OpInfo &matched_op_info) const {649 OpInfo &matched_op_info) const {
650+ if (IsHostCpuCustomOp(op_desc)) {
651+ SetHostCpuCustomOp(op_desc, matched_op_info);
652+ return kEngineNameCustom;
653+ }
613 for (const auto &it : op_infos) {654 for (const auto &it : op_infos) {
614 if ((it.engine == kHostCpuEngineName) && (it.opKernelLib == kHostCpuOpKernelLibName)) {655 if ((it.engine == kHostCpuEngineName) && (it.opKernelLib == kHostCpuOpKernelLibName)) {
615 op_desc->SetOpEngineName(kHostCpuEngineName);656 op_desc->SetOpEngineName(kHostCpuEngineName);
@@ -20,6 +20,8 @@
20#include "framework/common/debug/log.h"20#include "framework/common/debug/log.h"
21#include "framework/common/framework_types_internal.h"21#include "framework/common/framework_types_internal.h"
22#include "graph/debug/ge_attr_define.h"22#include "graph/debug/ge_attr_define.h"
23+#include "graph/custom_op_factory.h"
24+#include "graph/ascend_string.h"
23#include "graph/utils/graph_utils.h"25#include "graph/utils/graph_utils.h"
24#include "graph/utils/op_desc_utils.h"26#include "graph/utils/op_desc_utils.h"
25#include "graph/utils/node_utils.h"27#include "graph/utils/node_utils.h"
@@ -51,6 +53,17 @@ constexpr int64_t kThresholdForMergeAllToUnknownGraph = -1;
51constexpr int32_t kBase = 10;53constexpr int32_t kBase = 10;
52const std::string kStableRdfsSort = "3";54const std::string kStableRdfsSort = "3";
53constexpr char_t const *kOffline = "offline";55constexpr char_t const *kOffline = "offline";
56+ 
57+bool IsCustomOpExecOnHostCpu(const ge::OpDescPtr &op_desc) {
58+ if ((op_desc == nullptr) || (op_desc->GetOpEngineName() != ge::kEngineNameCustom) ||
59+ (op_desc->GetOpKernelLibName() != ge::kCustomOpKernelLibName) ||
60+ !ge::CustomOpFactory::IsExistOp(ge::AscendString(op_desc->GetTypePtr()), ge::OpBackend::kHostCPU)) {
61+ return false;
62+ }
63+ std::string lowering_func;
64+ return ge::AttrUtils::GetStr(op_desc, ge::kAttrLowingFunc, lowering_func) &&
65+ (lowering_func == ge::kHostCpuCustomOpLowerFunc);
66+}
54} // namespace67} // namespace
55 68 
56namespace ge {69namespace ge {
@@ -1258,7 +1271,8 @@ Status DynamicShapePartitioner::IsUnknownShapeNode(NodePtr node, bool &is_unknow
1258 }1271 }
1259 auto graph = GetRootGraph();1272 auto graph = GetRootGraph();
1260 GELOGD("node: %s, engine_name: %s.", node->GetNamePtr(), opdesc->GetOpEngineName().c_str());1273 GELOGD("node: %s, engine_name: %s.", node->GetNamePtr(), opdesc->GetOpEngineName().c_str());
1261- if (opdesc->GetOpEngineName() == kHostCpuEngineName) {1274+ 
1275+ if ((opdesc->GetOpEngineName() == kHostCpuEngineName) || IsCustomOpExecOnHostCpu(opdesc)) {
1262 is_unknown = true;1276 is_unknown = true;
1263 GELOGD("Mark host cpu node %s unknown as host engine as it relies on the runtime scheduler for execution.",1277 GELOGD("Mark host cpu node %s unknown as host engine as it relies on the runtime scheduler for execution.",
1264 node->GetName().c_str());1278 node->GetName().c_str());
@@ -1389,7 +1403,7 @@ Status DynamicShapePartitioner::CheckIfSubgraphUnknown(const ComputeGraphPtr &gr
1389 auto desc = node->GetOpDesc();1403 auto desc = node->GetOpDesc();
1390 GE_CHK_GRAPH_STATUS_RET(ge::NodeUtils::GetNodeUnknownShapeStatus(*node, is_unknown_shape),1404 GE_CHK_GRAPH_STATUS_RET(ge::NodeUtils::GetNodeUnknownShapeStatus(*node, is_unknown_shape),
1391 "[Get][ShapeStatus] of node[%s] failed!", node->GetName().c_str());1405 "[Get][ShapeStatus] of node[%s] failed!", node->GetName().c_str());
1392- if (desc->GetOpEngineName() == "DNN_VM_HOST_CPU") {1406+ if ((desc->GetOpEngineName() == kHostCpuEngineName) || IsCustomOpExecOnHostCpu(desc)) {
1393 is_unknown_shape = true;1407 is_unknown_shape = true;
1394 GELOGD("Mark host cpu node %s unknown.", node->GetName().c_str());1408 GELOGD("Mark host cpu node %s unknown.", node->GetName().c_str());
1395 }1409 }
@@ -20,10 +20,13 @@
20#include "common/plugin/ge_make_unique_util.h"20#include "common/plugin/ge_make_unique_util.h"
21#include "framework/common/op/ge_op_utils.h"21#include "framework/common/op/ge_op_utils.h"
22#include "common/compile_profiling/ge_trace_wrapper.h"22#include "common/compile_profiling/ge_trace_wrapper.h"
23+#include "graph/ascend_string.h"
24+#include "graph/custom_op_factory.h"
23#include "graph/utils/graph_utils.h"25#include "graph/utils/graph_utils.h"
24#include "graph/utils/op_desc_utils.h"26#include "graph/utils/op_desc_utils.h"
25#include "graph/utils/type_utils.h"27#include "graph/utils/type_utils.h"
26#include "graph/utils/op_type_utils.h"28#include "graph/utils/op_type_utils.h"
29+#include "graph/utils/attr_utils.h"
27#include "graph/build/stream/stream_utils.h"30#include "graph/build/stream/stream_utils.h"
28#include "common/checker.h"31#include "common/checker.h"
29#include "graph/ge_context.h"32#include "graph/ge_context.h"
@@ -48,11 +51,23 @@ const char_t *const kTaskL2FusionInfo = "_task_L2FusionInfo";
48const char_t *const kDataAnchorIndexForLxfusion = "_data_anchor_index_for_lxfusion";51const char_t *const kDataAnchorIndexForLxfusion = "_data_anchor_index_for_lxfusion";
49const char_t *const kEnableCvParallel = "_enable_cv_parallel";52const char_t *const kEnableCvParallel = "_enable_cv_parallel";
50const char_t *const kVectorEngineName = "VectorEngine";53const char_t *const kVectorEngineName = "VectorEngine";
54+const char_t *const kHostCpuEngineName = "DNN_VM_HOST_CPU";
51const std::string kStableRdfsSort = "3";55const std::string kStableRdfsSort = "3";
52const int32_t kOneGraph = 1; // only one graph56const int32_t kOneGraph = 1; // only one graph
53const int32_t kRankOne = 1; // order of graph list is 0,1,2,3..., 1 means second order57const int32_t kRankOne = 1; // order of graph list is 0,1,2,3..., 1 means second order
54const int32_t kRankZero = 0; // order of graph list is 0,1,2,3..., 0 means first order58const int32_t kRankZero = 0; // order of graph list is 0,1,2,3..., 0 means first order
55const int64_t kOverflowDefaultValue = -1;59const int64_t kOverflowDefaultValue = -1;
60+ 
61+bool IsCustomOpExecOnHostCpu(const OpDescPtr &op_desc) {
62+ if ((op_desc == nullptr) || (op_desc->GetOpEngineName() != kEngineNameCustom) ||
63+ (op_desc->GetOpKernelLibName() != kCustomOpKernelLibName) ||
64+ !CustomOpFactory::IsExistOp(AscendString(op_desc->GetTypePtr()), OpBackend::kHostCPU)) {
65+ return false;
66+ }
67+ std::string lowering_func;
68+ return AttrUtils::GetStr(op_desc, kAttrLowingFunc, lowering_func) && (lowering_func == kHostCpuCustomOpLowerFunc);
69+}
70+ 
56struct DeviceIndex {71struct DeviceIndex {
57 std::string engine_type;72 std::string engine_type;
58 std::vector<int32_t> indices;73 std::vector<int32_t> indices;
@@ -78,8 +93,9 @@ struct DeviceIndex {
78 93 
79std::string GenClusterEngineName(const NodePtr &node, EnginePartitioner::Mode mode, const NodeEngineMap &engine_map) {94std::string GenClusterEngineName(const NodePtr &node, EnginePartitioner::Mode mode, const NodeEngineMap &engine_map) {
80 auto engine_name = engine_map.at(node);95 auto engine_name = engine_map.at(node);
81- // 流分配时,自定义算子需要跟aicore算子一条流96+ // 流分配时,device自定义算子需要跟aicore算子一条流
82- if ((mode == EnginePartitioner::Mode::kSecondPartitioning) && (engine_name == kEngineNameCustom)) {97+ if ((mode == EnginePartitioner::Mode::kSecondPartitioning) && (engine_name == kEngineNameCustom) &&
98+ !IsCustomOpExecOnHostCpu(node->GetOpDesc())) {
83 // 临时改动,后续需要从用户的注册引擎信息里面获取此处的自定义算子挂靠的引擎名字99 // 临时改动,后续需要从用户的注册引擎信息里面获取此处的自定义算子挂靠的引擎名字
84 engine_name = kEngineNameAiCore;100 engine_name = kEngineNameAiCore;
85 }101 }
@@ -15,6 +15,7 @@
15#include "common/checker.h"15#include "common/checker.h"
16#include "mmpa/mmpa_api.h"16#include "mmpa/mmpa_api.h"
17#include "graph/ge_local_context.h"17#include "graph/ge_local_context.h"
18+#include "graph/custom_op_factory.h"
18#include "ge/ge_api_types.h"19#include "ge/ge_api_types.h"
19#include "common/math/ge_math_util.h"20#include "common/math/ge_math_util.h"
20#include "graph/option/optimization_option_info.h"21#include "graph/option/optimization_option_info.h"
@@ -75,13 +76,22 @@ bool IsGelocalOp(const OpDescPtr &op_desc) {
75 return op_desc->GetOpKernelLibName() == kGeLocalOpKernelLibName;76 return op_desc->GetOpKernelLibName() == kGeLocalOpKernelLibName;
76}77}
77 78 
79+bool IsDeviceCustomOp(const OpDescPtr &op_desc) {
80+ return CustomOpFactory::IsExistOp(AscendString(op_desc->GetTypePtr()), OpBackend::kDevice);
81+}
82+ 
83+bool IsHostCpuCustomOp(const OpDescPtr &op_desc) {
84+ return CustomOpFactory::IsExistOp(AscendString(op_desc->GetTypePtr()), OpBackend::kHostCPU);
85+}
86+ 
78bool IsControlV2Op(const std::string &op_type) {87bool IsControlV2Op(const std::string &op_type) {
79 return kControlV2Types.count(op_type) > 0U;88 return kControlV2Types.count(op_type) > 0U;
80}89}
81 90 
82bool IsExecOnDevice(const OpDescPtr &op_desc) {91bool IsExecOnDevice(const OpDescPtr &op_desc) {
83 return (op_desc->GetOpKernelLibName() == kEngineNameAiCpu) || (op_desc->GetOpKernelLibName() == kEngineNameAiCpuTf) ||92 return (op_desc->GetOpKernelLibName() == kEngineNameAiCpu) || (op_desc->GetOpKernelLibName() == kEngineNameAiCpuTf) ||
84- (op_desc->GetOpKernelLibName() == kEngineNameAiCore);93+ (op_desc->GetOpKernelLibName() == kEngineNameAiCore) ||
94+ (op_desc->GetOpKernelLibName() == kCustomOpKernelLibName && IsDeviceCustomOp(op_desc));
85}95}
86 96 
87bool IsConstOp(const OpDescPtr &op_desc) {97bool IsConstOp(const OpDescPtr &op_desc) {
@@ -303,6 +313,20 @@ bool HostcpuEngineUpdatePass::CheckAndMarkHostExec(const NodePtr &node, NodeEngi
303 return true;313 return true;
304 }314 }
305 315 
316+ if (IsHostCpuCustomOp(op_desc) && IsExecOnDevice(op_desc)) {
317+ op_desc->SetOpEngineName(kEngineNameCustom);
318+ op_desc->SetOpKernelLibName(kCustomOpKernelLibName);
319+ (void)AttrUtils::SetStr(op_desc, ATTR_NAME_ENGINE_NAME_FOR_LX, kEngineNameCustom);
320+ (void)AttrUtils::SetStr(op_desc, ATTR_NAME_KKERNEL_LIB_NAME_FOR_LX, kCustomOpKernelLibName);
321+ node_atomic_engine_map[node] = kEngineNameCustom;
322+ node_composite_engine_map[node] = kEngineNameCustom;
323+ (void)AttrUtils::SetStr(op_desc, kAttrLowingFunc, kHostCpuCustomOpLowerFunc);
324+ GELOGI("[HostcpuEngineUpdatePass]: Set OpKernelLibName %s and OpEngineName %s to %s",
325+ kCustomOpKernelLibName.c_str(), kEngineNameCustom.c_str(), op_desc->GetName().c_str());
326+ host_exe_ops_.insert(node);
327+ return true;
328+ }
329+ 
306 if (IsSupportHostcpu(op_desc) && IsExecOnDevice(op_desc)) {330 if (IsSupportHostcpu(op_desc) && IsExecOnDevice(op_desc)) {
307 op_desc->SetOpEngineName(kHostCpuEngineName);331 op_desc->SetOpEngineName(kHostCpuEngineName);
308 op_desc->SetOpKernelLibName(kHostCpuOpKernelLibName);332 op_desc->SetOpKernelLibName(kHostCpuOpKernelLibName);
@@ -12,11 +12,24 @@
12#include "graph/utils/node_utils.h"12#include "graph/utils/node_utils.h"
13#include "graph/debug/ge_attr_define.h"13#include "graph/debug/ge_attr_define.h"
14#include "graph/ge_context.h"14#include "graph/ge_context.h"
15+#include "graph/custom_op_factory.h"
16+#include "graph/ascend_string.h"
17+#include "common/ge_common/ge_types.h"
15 18 
16namespace ge {19namespace ge {
17namespace {20namespace {
18const char *const kOwnerGraphIsUnknown = "OwnerGraphIsUnknown";21const char *const kOwnerGraphIsUnknown = "OwnerGraphIsUnknown";
19const char *const kHostCpuEngineName = "DNN_VM_HOST_CPU";22const char *const kHostCpuEngineName = "DNN_VM_HOST_CPU";
23+ 
24+bool IsCustomOpExecOnHostCpu(const OpDescPtr &op_desc) {
25+ if ((op_desc == nullptr) || (op_desc->GetOpEngineName() != kEngineNameCustom) ||
26+ (op_desc->GetOpKernelLibName() != kCustomOpKernelLibName) ||
27+ !CustomOpFactory::IsExistOp(AscendString(op_desc->GetTypePtr()), OpBackend::kHostCPU)) {
28+ return false;
29+ }
30+ std::string lowering_func;
31+ return AttrUtils::GetStr(op_desc, kAttrLowingFunc, lowering_func) && (lowering_func == kHostCpuCustomOpLowerFunc);
32+}
20} // namespace33} // namespace
21 34 
22Status MarkGraphUnknownStatusPass::Run(ComputeGraphPtr graph) {35Status MarkGraphUnknownStatusPass::Run(ComputeGraphPtr graph) {
@@ -36,7 +49,7 @@ Status MarkGraphUnknownStatusPass::Run(ComputeGraphPtr graph) {
36 }49 }
37 GE_CHK_GRAPH_STATUS_RET(ge::NodeUtils::GetNodeUnknownShapeStatus(*node, is_unknown_shape),50 GE_CHK_GRAPH_STATUS_RET(ge::NodeUtils::GetNodeUnknownShapeStatus(*node, is_unknown_shape),
38 "[Get][ShapeStatus] of node[%s] failed!", node->GetName().c_str());51 "[Get][ShapeStatus] of node[%s] failed!", node->GetName().c_str());
39- if (desc->GetOpEngineName() == kHostCpuEngineName) {52+ if ((desc->GetOpEngineName() == kHostCpuEngineName) || IsCustomOpExecOnHostCpu(desc)) {
40 is_unknown_shape = true;53 is_unknown_shape = true;
41 GELOGD("Mark host cpu node %s unknown.", node->GetName().c_str());54 GELOGD("Mark host cpu node %s unknown.", node->GetName().c_str());
42 }55 }
@@ -10,9 +10,23 @@
10 10 
11#include "graph/passes/standard_optimize/constant_folding/constant_folding_pass.h"11#include "graph/passes/standard_optimize/constant_folding/constant_folding_pass.h"
12 12 
13+#include <memory>
14+#include <new>
15+#include <utility>
13#include <vector>16#include <vector>
17+#include "graph_metadef/common/ge_common/util.h"
18+#include "rt_external_mem.h"
19+#include "common/memory/tensor_trans_utils.h"
20+#include "exe_graph/runtime/host_cpu_op_execution_context.h"
21+#include "exe_graph/runtime/gert_mem_allocator.h"
22+#include "exe_graph/runtime/runtime_tensor.h"
23+#include "exe_graph/lowering/kernel_run_context_builder.h"
24+#include "graph/custom_op.h"
25+#include "graph/custom_op/cast.h"
26+#include "graph/custom_op_factory.h"
27+#include "graph/ge_tensor.h"
28+#include "graph/op_desc.h"
14#include "graph/utils/node_utils.h"29#include "graph/utils/node_utils.h"
15-#include "graph/utils/type_utils.h"
16#include "graph/utils/constant_utils.h"30#include "graph/utils/constant_utils.h"
17#include "host_cpu_engine/host_cpu_engine.h"31#include "host_cpu_engine/host_cpu_engine.h"
18#include "api/gelib/gelib.h"32#include "api/gelib/gelib.h"
@@ -26,6 +40,155 @@ const int64_t kShapeCalNum = 8;
26const char *const kKernelLibName = "aicpu_ascend_kernel";40const char *const kKernelLibName = "aicpu_ascend_kernel";
27const char *const kOpsFlagClose = "0";41const char *const kOpsFlagClose = "0";
28const char *const kPassName = "ConstantFoldingPass";42const char *const kPassName = "ConstantFoldingPass";
43+ 
44+class HostCpuConstFoldingMemAllocator final : public ge::Allocator {
45+ public:
46+ ge::MemBlock *Malloc(size_t size) override {
47+ auto *buffer = new (std::nothrow) uint8_t[size];
48+ if (buffer == nullptr) {
49+ return nullptr;
50+ }
51+ auto *block = new (std::nothrow) ge::MemBlock(*this, buffer, size);
52+ if (block == nullptr) {
53+ delete[] buffer;
54+ return nullptr;
55+ }
56+ return block;
57+ }
58+ 
59+ void Free(ge::MemBlock *block) override {
60+ if (block != nullptr) {
61+ delete[] static_cast<uint8_t *>(block->GetAddr());
62+ delete block;
63+ }
64+ }
65+};
66+ 
67+HostCpuConstFoldingMemAllocator &GetHostCpuConstFoldingMemAllocator() {
68+ static HostCpuConstFoldingMemAllocator allocator;
69+ return allocator;
70+}
71+ 
72+ge::graphStatus HostCpuMemBlockManager(void *block, gert::TensorOperateType operate_type, void **out) {
73+ GE_ASSERT_NOTNULL(block);
74+ auto *mem_block = reinterpret_cast<ge::MemBlock *>(block);
75+ GE_ASSERT((operate_type == gert::kGetTensorAddress || operate_type == gert::kFreeTensor ||
76+ operate_type == gert::kPlusShareCount),
77+ "Unexpected operate type %d", static_cast<int32_t>(operate_type));
78+ if (operate_type == gert::kGetTensorAddress) {
79+ GE_ASSERT_NOTNULL(out);
80+ *out = mem_block->GetAddr();
81+ }
82+ if (operate_type == gert::kPlusShareCount) {
83+ mem_block->AddCount();
84+ }
85+ if (operate_type == gert::kFreeTensor) {
86+ mem_block->Free();
87+ }
88+ return ge::GRAPH_SUCCESS;
89+}
90+ 
91+gert::TensorData MakeHostCpuConstFoldingTensorData(size_t size, gert::TensorPlacement placement) {
92+ auto *mem_block = GetHostCpuConstFoldingMemAllocator().Malloc(size);
93+ if ((mem_block == nullptr) || (mem_block->GetAddr() == nullptr)) {
94+ if (mem_block != nullptr) {
95+ GetHostCpuConstFoldingMemAllocator().Free(mem_block);
96+ }
97+ return gert::TensorData();
98+ }
99+ return gert::TensorData(mem_block, HostCpuMemBlockManager, size, placement);
100+}
101+ 
102+class HostCpuConstFoldingMemGertAllocator final : public gert::GertAllocator {
103+ public:
104+ HostCpuConstFoldingMemGertAllocator() : gert::GertAllocator(0, gert::kOnHost) {}
105+ ~HostCpuConstFoldingMemGertAllocator() override = default;
106+ 
107+ gert::GertMemBlock *Malloc(size_t size) override {
108+ (void)size;
109+ return nullptr;
110+ }
111+ 
112+ gert::GertTensorData MallocTensorData(size_t size) override {
113+ const auto tensor_data = MakeHostCpuConstFoldingTensorData(size, GetPlacement());
114+ if (tensor_data.GetAddr() == nullptr) {
115+ return {};
116+ }
117+ gert::GertTensorData gtd;
118+ if (gtd.MutableTensorData().ShareFrom(tensor_data) != ge::GRAPH_SUCCESS) {
119+ return {};
120+ }
121+ return gtd;
122+ }
123+ 
124+ gert::TensorData MallocTensorDataFromL1(size_t size) override {
125+ return MakeHostCpuConstFoldingTensorData(size, GetPlacement());
126+ }
127+ 
128+ void Free(gert::GertMemBlock *block) override {
129+ (void)block;
130+ }
131+ 
132+ ge::graphStatus FreeAt(int64_t stream_id, gert::GertMemBlock *block) override {
133+ (void)stream_id;
134+ (void)block;
135+ return ge::GRAPH_SUCCESS;
136+ }
137+ 
138+ ge::graphStatus ShareFromTensorData(const gert::TensorData &td, gert::GertTensorData &gtd) override {
139+ return gtd.MutableTensorData().ShareFrom(td);
140+ }
141+ 
142+ int64_t GetStreamNum() override {
143+ return 1;
144+ }
145+ 
146+ ge::graphStatus SetL1Allocator(ge::Allocator *allocator) override {
147+ (void)allocator;
148+ return ge::GRAPH_SUCCESS;
149+ }
150+};
151+ 
152+Status BuildHostCpuOpContext(const NodePtr &node, const std::vector<ConstGeTensorPtr> &inputs,
153+ std::vector<gert::Tensor> &input_tensors, std::vector<gert::Tensor> &output_tensors,
154+ HostCpuConstFoldingMemGertAllocator &allocator,
155+ gert::KernelContextHolder &context_holder) {
156+ const auto op_desc = node->GetOpDesc();
157+ GE_ASSERT_NOTNULL(op_desc);
158+ const size_t input_num = inputs.size();
159+ const size_t output_num = op_desc->GetOutputsSize();
160+ input_tensors.resize(input_num);
161+ output_tensors.resize(output_num);
162+ 
163+ for (size_t i = 0U; i < input_num; ++i) {
164+ GE_ASSERT_SUCCESS(TensorTransUtils::GeTensor2GertTensor(*inputs[i], input_tensors[i]));
165+ }
166+ 
167+ for (size_t i = 0U; i < output_num; ++i) {
168+ GE_ASSERT_SUCCESS(TensorTransUtils::GeTensor2GertTensor(GeTensor(op_desc->GetOutputDesc(i)), output_tensors[i]));
169+ }
170+ 
171+ std::vector<void *> input_ptrs;
172+ input_ptrs.reserve(input_num + 1U);
173+ for (auto &input_tensor : input_tensors) {
174+ input_ptrs.emplace_back(&input_tensor);
175+ }
176+ input_ptrs.emplace_back(&allocator);
177+ 
178+ std::vector<void *> output_ptrs;
179+ output_ptrs.reserve(output_num);
180+ for (auto &output_tensor : output_tensors) {
181+ output_ptrs.emplace_back(&output_tensor);
182+ }
183+ 
184+ context_holder =
185+ gert::KernelRunContextBuilder().Inputs(std::move(input_ptrs)).Outputs(std::move(output_ptrs)).Build(op_desc);
186+ if (context_holder.GetKernelContext() == nullptr) {
187+ GELOGE(FAILED, "Build HostCpu op context failed for node %s.", node->GetName().c_str());
188+ return FAILED;
189+ }
190+ return SUCCESS;
191+}
29} // namespace192} // namespace
30 193 
31bool ConstantFoldingPass::NeedIgnorePass(const NodePtr &node) {194bool ConstantFoldingPass::NeedIgnorePass(const NodePtr &node) {
@@ -82,19 +245,62 @@ Status ConstantFoldingPass::ComputePotentialWeight(NodePtr &node, std::vector<Ge
82 }245 }
83 // Try to run kernel on host cpu246 // Try to run kernel on host cpu
84 uint64_t start_time = GetCurrentTimestamp();247 uint64_t start_time = GetCurrentTimestamp();
85- Status compute_ret = ComputeWithHostCpuKernel(node, inputs, outputs);248+ Status compute_ret = ComputeWithHostCpuCustomOp(node, inputs, outputs);
86 if (compute_ret == SUCCESS) {249 if (compute_ret == SUCCESS) {
87 CollectCostTimeOfOpConstantFolding(node, start_time);250 CollectCostTimeOfOpConstantFolding(node, start_time);
88 } else {251 } else {
89- // If computation on AICPU is not possible, try running the host kernel within GE.252+ // If host custom op computation is not possible, try running the HostCpu kernel.
90- GELOGD("Try to compute weight of %s with built-in kernel.", node->GetName().c_str());253+ GELOGD("Try to compute weight of %s with HostCpu kernel.", node->GetName().c_str());
91- compute_ret = ComputeWithBuiltInKernel(node, inputs, outputs);254+ compute_ret = ComputeWithHostCpuKernel(node, inputs, outputs);
255+ if (compute_ret == SUCCESS) {
256+ CollectCostTimeOfOpConstantFolding(node, start_time);
257+ } else {
258+ // If computation on AICPU is not possible, try running the host kernel within GE.
259+ GELOGD("Try to compute weight of %s with built-in kernel.", node->GetName().c_str());
260+ compute_ret = ComputeWithBuiltInKernel(node, inputs, outputs);
261+ }
92 }262 }
93 GELOGD("Constant folding computation for node %s (type: %s) finished, return code: %u.", node->GetName().c_str(),263 GELOGD("Constant folding computation for node %s (type: %s) finished, return code: %u.", node->GetName().c_str(),
94 node->GetType().c_str(), compute_ret);264 node->GetType().c_str(), compute_ret);
95 return compute_ret;265 return compute_ret;
96}266}
97 267 
268+Status ConstantFoldingPass::ComputeWithHostCpuCustomOp(const NodePtr &node, const vector<ConstGeTensorPtr> &inputs,
269+ std::vector<GeTensorPtr> &outputs) {
270+ const std::string op_type = NodeUtils::GetNodeType(node);
271+ const AscendString op_type_str(op_type.c_str());
272+ if (!CustomOpFactory::IsExistOp(op_type_str, OpBackend::kHostCPU)) {
273+ GELOGD("Op of type %s is not supported by host cpu custom op.", op_type.c_str());
274+ return UNSUPPORTED;
275+ }
276+ auto base_custom_op = CustomOpFactory::CreateOrGetCustomOp(op_type_str, OpBackend::kHostCPU);
277+ GE_ASSERT_NOTNULL(base_custom_op, "Op %s is registered as host cpu custom op but create instance failed.",
278+ op_type.c_str());
279+ auto *host_custom_op = CustomOpCast<HostCpuExecuteOp>(base_custom_op);
280+ GE_ASSERT_NOTNULL(host_custom_op,
281+ "Op %s is registered as host cpu custom op but does not implement HostCpuExecuteOp.",
282+ op_type.c_str());
283+ 
284+ std::vector<gert::Tensor> input_tensors;
285+ std::vector<gert::Tensor> output_tensors;
286+ HostCpuConstFoldingMemGertAllocator allocator;
287+ gert::KernelContextHolder context_holder;
288+ GE_ASSERT_SUCCESS(BuildHostCpuOpContext(node, inputs, input_tensors, output_tensors, allocator, context_holder));
289+ auto *host_context = reinterpret_cast<gert::HostCpuOpExecutionContext *>(context_holder.GetKernelContext());
290+ GE_ASSERT_NOTNULL(host_context);
291+ GE_ASSERT_SUCCESS(host_custom_op->Execute(host_context));
292+ 
293+ outputs.clear();
294+ outputs.reserve(output_tensors.size());
295+ for (const auto &output_tensor : output_tensors) {
296+ GeTensorPtr output = MakeShared<GeTensor>();
297+ GE_ASSERT_NOTNULL(output);
298+ GE_ASSERT_SUCCESS(TensorTransUtils::GertTensor2GeTensor(output_tensor, *output));
299+ outputs.emplace_back(std::move(output));
300+ }
301+ return SUCCESS;
302+}
303+ 
98Status ConstantFoldingPass::ComputeWithBuiltInKernel(NodePtr &node, const vector<ConstGeTensorPtr> &inputs,304Status ConstantFoldingPass::ComputeWithBuiltInKernel(NodePtr &node, const vector<ConstGeTensorPtr> &inputs,
99 std::vector<GeTensorPtr> &outputs) {305 std::vector<GeTensorPtr> &outputs) {
100 auto op_kernel = folding_pass::GetKernelByType(node);306 auto op_kernel = folding_pass::GetKernelByType(node);
@@ -30,6 +30,8 @@ class ConstantFoldingPass : public PotentialFoldingPass {
30 30 
31 static Status RunOpKernel(const NodePtr &node, const std::vector<ConstGeTensorPtr> &inputs,31 static Status RunOpKernel(const NodePtr &node, const std::vector<ConstGeTensorPtr> &inputs,
32 std::vector<GeTensorPtr> &outputs);32 std::vector<GeTensorPtr> &outputs);
33+ static Status ComputeWithHostCpuCustomOp(const NodePtr &node, const std::vector<ConstGeTensorPtr> &inputs,
34+ std::vector<GeTensorPtr> &outputs);
33 static Status ComputeWithHostCpuKernel(const NodePtr &node, const std::vector<ConstGeTensorPtr> &inputs,35 static Status ComputeWithHostCpuKernel(const NodePtr &node, const std::vector<ConstGeTensorPtr> &inputs,
34 std::vector<GeTensorPtr> &outputs);36 std::vector<GeTensorPtr> &outputs);
35 37 
@@ -78,7 +78,7 @@ Tensor *HostCpuOpExecutionContext::MallocOutputTensor(size_t index, const Storag
78 return output_tensor;78 return output_tensor;
79}79}
80 80 
81-Tensor *HostCpuOpExecutionContext::MakeOutputRefInput(size_t output_index, size_t input_index) const {81+Tensor *HostCpuOpExecutionContext::MakeOutputRefInput(size_t output_index, size_t input_index) {
82 const auto additional_start_index = GetAdditionalInputStartIndex();82 const auto additional_start_index = GetAdditionalInputStartIndex();
83 GE_ASSERT_TRUE(additional_start_index >= 0);83 GE_ASSERT_TRUE(additional_start_index >= 0);
84 84 
@@ -88,7 +88,7 @@ Tensor *HostCpuOpExecutionContext::MakeOutputRefInput(size_t output_index, size_
88 auto output_name = op_desc->GetOutputNameByIndex(output_index);88 auto output_name = op_desc->GetOutputNameByIndex(output_index);
89 GE_ASSERT_TRUE(input_name == output_name, "[MakeOutputRefInput] output name does not exist in input");89 GE_ASSERT_TRUE(input_name == output_name, "[MakeOutputRefInput] output name does not exist in input");
90 }90 }
91- auto *output_tensor = const_cast<Tensor *>(GetOutputPointer<Tensor>(output_index));91+ auto *output_tensor = GetOutputPointer<Tensor>(output_index);
92 GE_ASSERT_NOTNULL(output_tensor);92 GE_ASSERT_NOTNULL(output_tensor);
93 93 
94 auto input_tensor = GetInputPointer<Tensor>(input_index);94 auto input_tensor = GetInputPointer<Tensor>(input_index);
@@ -230,7 +230,7 @@ Tensor *HostCpuOpExecutionContext::MallocOutputTensor(size_t index, const Storag
230 return nullptr;230 return nullptr;
231}231}
232 232 
233-Tensor *HostCpuOpExecutionContext::MakeOutputRefInput(size_t output_index, size_t input_index) const {233+Tensor *HostCpuOpExecutionContext::MakeOutputRefInput(size_t output_index, size_t input_index) {
234 (void)output_index;234 (void)output_index;
235 (void)input_index;235 (void)input_index;
236 return nullptr;236 return nullptr;
@@ -105,6 +105,7 @@ const std::string kFFTSAiCoreLowerFunc = "ffts_ai_core_lower_func";
105const std::string kFFTSGraphLowerFunc = "ffts_graph_lower_func";105const std::string kFFTSGraphLowerFunc = "ffts_graph_lower_func";
106const std::string kFFTSStaticGraphLowerFunc = "ffts_static_graph_lower_func";106const std::string kFFTSStaticGraphLowerFunc = "ffts_static_graph_lower_func";
107const std::string kFFTSMixL2LowerFunc = "ffts_mix_l2_lower_func";107const std::string kFFTSMixL2LowerFunc = "ffts_mix_l2_lower_func";
108+const std::string kHostCpuCustomOpLowerFunc = "host_cpu_custom_op_lower_func";
108// runtime2.0 calculate func109// runtime2.0 calculate func
109const std::string kAttrCalcArgsSizeFunc = "_ge_attr_calculate_func";110const std::string kAttrCalcArgsSizeFunc = "_ge_attr_calculate_func";
110const std::string kFFTSMixL2CalcFunc = "ffts_mix_l2_calc_func";111const std::string kFFTSMixL2CalcFunc = "ffts_mix_l2_calc_func";
@@ -78,12 +78,12 @@ class HostCpuOpExecutionContext : public ExtendedKernelContext {
78 Tensor *MallocOutputTensor(size_t index, const StorageShape &shape, const StorageFormat &format, ge::DataType dtype);78 Tensor *MallocOutputTensor(size_t index, const StorageShape &shape, const StorageFormat &format, ge::DataType dtype);
79 79 
80 /**80 /**
81- * 指定某输出的内存地址引用自某个输入。81+ * 指定某输出的内存地址引用自某个输入,同时初始化tensor的基本信息。
82 * @param output_index 输出 index82 * @param output_index 输出 index
83 * @param input_index 输入 index83 * @param input_index 输入 index
84 * @return output_index 对应的输出 Tensor 指针,异常时返回空指针84 * @return output_index 对应的输出 Tensor 指针,异常时返回空指针
85 */85 */
86- Tensor *MakeOutputRefInput(size_t output_index, size_t input_index) const;86+ Tensor *MakeOutputRefInput(size_t output_index, size_t input_index);
87 87 
88 enum class AdditionalInputIndex : uint32_t { kHostAllocator = 0U, kNum };88 enum class AdditionalInputIndex : uint32_t { kHostAllocator = 0U, kNum };
89 89 
@@ -85,6 +85,18 @@ bg::ValueHolderPtr FindCustomExecutorFunc(const ge::NodePtr &node, const LowerIn
85 return lower_input.global_data->GetOrCreateUniqueValueHolder(node->GetType() + "_FindCustomOp_", builder)[0];85 return lower_input.global_data->GetOrCreateUniqueValueHolder(node->GetType() + "_FindCustomOp_", builder)[0];
86}86}
87 87 
88+bg::ValueHolderPtr FindHostCpuCustomExecutorFunc(const ge::NodePtr &node, const LowerInput &lower_input) {
89+ auto builder = [&node, &lower_input]() -> std::vector<bg::ValueHolderPtr> {
90+ return bg::FrameSelector::OnInitRoot([&node, &lower_input]() -> std::vector<bg::ValueHolderPtr> {
91+ auto node_type = bg::ValueHolder::CreateConst(node->GetTypePtr(), node->GetType().size() + 1, true);
92+ ge::CustomOpRegistry *custom_op_registry = lower_input.global_data->GetCustomOpRegistry().get();
93+ auto registry_holder = bg::ValueHolder::CreateConst(&custom_op_registry, sizeof(custom_op_registry));
94+ return {bg::ValueHolder::CreateSingleDataOutput("FindHostCpuCustomOp", {node_type, registry_holder})};
95+ });
96+ };
97+ return lower_input.global_data->GetOrCreateUniqueValueHolder(node->GetType() + "_FindHostCpuCustomOp_", builder)[0];
98+}
99+ 
88ge::graphStatus BuildInputTensors(const ge::NodePtr &node, const LowerInput &lower_input,100ge::graphStatus BuildInputTensors(const ge::NodePtr &node, const LowerInput &lower_input,
89 std::vector<bg::ValueHolderPtr> &input_tensor_holders,101 std::vector<bg::ValueHolderPtr> &input_tensor_holders,
90 std::vector<bg::ValueHolderPtr> &input_addr_holders) {102 std::vector<bg::ValueHolderPtr> &input_addr_holders) {
@@ -120,6 +132,37 @@ ge::graphStatus BuildInputTensors(const ge::NodePtr &node, const LowerInput &low
120 lower_input.input_shapes.size(), input_tensor_holders.size());132 lower_input.input_shapes.size(), input_tensor_holders.size());
121 return ge::SUCCESS;133 return ge::SUCCESS;
122}134}
135+ 
136+ge::graphStatus BuildHostInputTensors(const ge::NodePtr &node, const LowerInput &lower_input,
137+ std::vector<bg::ValueHolderPtr> &input_tensor_holders,
138+ std::vector<bg::ValueHolderPtr> &input_addr_holders) {
139+ for (const ge::InDataAnchorPtr &in_data_anchor : node->GetAllInDataAnchors()) {
140+ GE_ASSERT_NOTNULL(in_data_anchor);
141+ // optional场景
142+ ge::OutDataAnchorPtr out_data_anchor = in_data_anchor->GetPeerOutAnchor();
143+ if (out_data_anchor == nullptr) {
144+ continue;
145+ }
146+ ge::NodePtr peer_node = out_data_anchor->GetOwnerNode();
147+ GE_ASSERT_NOTNULL(peer_node);
148+ const auto *const_lower_result = peer_node->GetOpDesc()->GetExtAttr<PlacedLoweringResult>(kLoweringResult);
149+ GE_ASSERT_NOTNULL(const_lower_result, "Lowering result of node [%s, %s] is not found.", peer_node->GetNamePtr(),
150+ peer_node->GetTypePtr());
151+ auto *lower_result = const_cast<PlacedLoweringResult *>(const_lower_result);
152+ GE_ASSERT_NOTNULL(lower_result);
153+ const OutputLowerResult *result = lower_result->GetOutputTensorResult(
154+ *lower_input.global_data, out_data_anchor->GetIdx(), {kOnHost, node->GetOpDesc()->GetStreamId()});
155+ GE_ASSERT_NOTNULL(result, "Lowering result of node [%s, %s] output[%d] is nullptr.", peer_node->GetNamePtr(),
156+ peer_node->GetTypePtr(), out_data_anchor->GetIdx());
157+ GE_ASSERT_NOTNULL(result->shape);
158+ input_tensor_holders.emplace_back(result->shape);
159+ input_addr_holders.emplace_back(result->address);
160+ }
161+ GE_ASSERT_TRUE(lower_input.input_shapes.size() == input_tensor_holders.size(),
162+ "Size[%zu] of input shapes and size[%zu] of input tensor is not same.",
163+ lower_input.input_shapes.size(), input_tensor_holders.size());
164+ return ge::SUCCESS;
165+}
123} // namespace166} // namespace
124 167 
125LowerResult LoweringCustomNode(const ge::NodePtr &node, const LowerInput &lower_input) {168LowerResult LoweringCustomNode(const ge::NodePtr &node, const LowerInput &lower_input) {
@@ -179,4 +222,51 @@ LowerResult LoweringCustomNode(const ge::NodePtr &node, const LowerInput &lower_
179 return {HyperStatus::Success(), {}, output_shapes, output_addrs};222 return {HyperStatus::Success(), {}, output_shapes, output_addrs};
180}223}
181REGISTER_NODE_CONVERTER_PLACEMENT(ge::kCustomOpKernelLibName.c_str(), kOnDeviceHbm, LoweringCustomNode);224REGISTER_NODE_CONVERTER_PLACEMENT(ge::kCustomOpKernelLibName.c_str(), kOnDeviceHbm, LoweringCustomNode);
225+ 
226+LowerResult LoweringHostCustomNode(const ge::NodePtr &node, const LowerInput &lower_input) {
227+ LOWER_REQUIRE_HYPER_SUCCESS(CheckLowerInput(lower_input));
228+ std::vector<bg::ValueHolderPtr> input_holders;
229+ std::vector<bg::ValueHolderPtr> input_addr_holders;
230+ LOWER_REQUIRE_SUCCESS(BuildHostInputTensors(node, lower_input, input_holders, input_addr_holders));
231+ // Allocate
232+ auto allocator_holder = lower_input.global_data->GetOrCreateAllocator({kOnHost, AllocatorUsage::kAllocNodeOutput});
233+ // Create op executeFunc
234+ auto custom_executor_func = FindHostCpuCustomExecutorFunc(node, lower_input);
235+ input_holders.emplace_back(allocator_holder);
236+ input_holders.emplace_back(custom_executor_func);
237+ // Check inference_rule
238+ const auto op_desc = node->GetOpDesc();
239+ LOWER_REQUIRE_NOTNULL(op_desc);
240+ std::string kernel_type = "ExecuteHostCustomOp";
241+ if (NeedCustomOpInferShape(node, *lower_input.global_data)) {
242+ kernel_type = "ExecuteHostCustomOpWithInferShape";
243+ std::vector<bg::ValueHolderPtr> infer_output_shapes =
244+ bg::InferCustomOpShape(node, lower_input.input_shapes, *lower_input.global_data);
245+ input_holders.insert(input_holders.end(), infer_output_shapes.begin(), infer_output_shapes.end());
246+ }
247+ std::vector<bg::ValueHolderPtr> output_tensor_holders =
248+ bg::ValueHolder::CreateDataOutput(kernel_type.c_str(), input_holders, node->GetAllOutDataAnchorsSize());
249+ std::vector<bg::ValueHolderPtr> output_shapes;
250+ std::vector<bg::DevMemValueHolderPtr> output_addrs;
251+ for (size_t i = 0UL; i < node->GetAllOutDataAnchorsSize(); i++) {
252+ auto split_outputs = bg::DevMemValueHolder::CreateDataOutput(
253+ kernel::kSplitDataTensor, {output_tensor_holders[i], allocator_holder},
254+ static_cast<size_t>(kernel::SplitTensorOutputs::kNum), op_desc->GetStreamId());
255+ CONVERTER_CHECK_HOLDERS_ALL_OK(split_outputs, static_cast<size_t>(kernel::SplitTensorOutputs::kNum));
256+ auto output_addr = split_outputs[static_cast<size_t>(kernel::SplitTensorOutputs::kTensorData)];
257+ output_addr->SetPlacement(kOnHost);
258+ LOWER_REQUIRE_NOTNULL(bg::ValueHolder::CreateVoidGuarder("FreeMemory", output_addr, {}));
259+ output_shapes.emplace_back(split_outputs[static_cast<size_t>(kernel::SplitTensorOutputs::kShape)]);
260+ output_addrs.emplace_back(output_addr);
261+ }
262+ 
263+ for (auto &addr : input_addr_holders) {
264+ auto guarder = addr->GetGuarder();
265+ if ((guarder != nullptr) && (!output_tensor_holders.empty())) {
266+ GE_ASSERT_HYPER_SUCCESS(bg::ValueHolder::AddDependency(output_tensor_holders.front(), guarder));
267+ }
268+ }
269+ return {HyperStatus::Success(), {}, output_shapes, output_addrs};
270+}
271+REGISTER_NODE_CONVERTER_PLACEMENT(ge::kHostCpuCustomOpLowerFunc.c_str(), kOnHost, LoweringHostCustomNode);
182} // namespace gert272} // namespace gert
@@ -15,5 +15,6 @@
15 15 
16namespace gert {16namespace gert {
17LowerResult LoweringCustomNode(const ge::NodePtr &node, const LowerInput &lower_input);17LowerResult LoweringCustomNode(const ge::NodePtr &node, const LowerInput &lower_input);
18-}18+LowerResult LoweringHostCustomNode(const ge::NodePtr &node, const LowerInput &lower_input);
19+} // namespace gert
19#endif // AIR_CXX_RUNTIME_V2_NODE_CUSTOM_CONVERTER_CUSTOM_NODE_CONVERTER_H_20#endif // AIR_CXX_RUNTIME_V2_NODE_CUSTOM_CONVERTER_CUSTOM_NODE_CONVERTER_H_
@@ -20,6 +20,7 @@
20#include "graph/def_types.h"20#include "graph/def_types.h"
21#include "graph/utils/type_utils.h"21#include "graph/utils/type_utils.h"
22#include "exe_graph/runtime/eager_op_execution_context.h"22#include "exe_graph/runtime/eager_op_execution_context.h"
23+#include "exe_graph/runtime/host_cpu_op_execution_context.h"
23#include "rt_external_kernel.h"24#include "rt_external_kernel.h"
24#include "core/executor/multi_thread_topological/executor/schedule/producer/producers/kernel_tags/critical_section_config.h"25#include "core/executor/multi_thread_topological/executor/schedule/producer/producers/kernel_tags/critical_section_config.h"
25#include "runtime/v2/engine/custom/kernel/eager_args_handler.h"26#include "runtime/v2/engine/custom/kernel/eager_args_handler.h"
@@ -29,6 +30,10 @@ namespace kernel {
29namespace {30namespace {
30// 自定义算子特有的输入,从 AdditionalInputIndex::kNum 开始31// 自定义算子特有的输入,从 AdditionalInputIndex::kNum 开始
31enum class CustomOpInput { kFunc = static_cast<uint32_t>(EagerOpExecutionContext::AdditionalInputIndex::kNum), kEnd };32enum class CustomOpInput { kFunc = static_cast<uint32_t>(EagerOpExecutionContext::AdditionalInputIndex::kNum), kEnd };
33+enum class HostCustomOpInput {
34+ kFunc = static_cast<uint32_t>(HostCpuOpExecutionContext::AdditionalInputIndex::kNum),
35+ kEnd
36+};
32 37 
33std::string PrintNodeType(const KernelContext *context) {38std::string PrintNodeType(const KernelContext *context) {
34 std::stringstream ss;39 std::stringstream ss;
@@ -90,6 +95,37 @@ ge::graphStatus FindCustomOpFunc(KernelContext *context) {
90 return ge::GRAPH_SUCCESS;95 return ge::GRAPH_SUCCESS;
91}96}
92 97 
98+ge::graphStatus FindHostCpuCustomOpFunc(KernelContext *context) {
99+ const char *node_type = context->GetInputValue<char *>(0);
100+ GE_ASSERT_NOTNULL(node_type, "Failed to find host CPU custom op func, node type is nullptr");
101+ auto custom_op_registry = context->GetInputValue<ge::CustomOpRegistry *>(1);
102+ GE_ASSERT_NOTNULL(custom_op_registry, "Failed to find host CPU custom op func, custom op registry is nullptr.");
103+ ge::BaseCustomOp *custom_op_ptr = custom_op_registry->CreateOrGetCustomOp(node_type, ge::OpBackend::kHostCPU);
104+ GE_ASSERT_NOTNULL(custom_op_ptr, "Failed to find host CPU custom op func for op type %s in custom op registry.",
105+ node_type);
106+ auto chain = context->GetOutput(0);
107+ GE_ASSERT_NOTNULL(chain);
108+ chain->Set(custom_op_ptr, nullptr);
109+ return ge::GRAPH_SUCCESS;
110+}
111+ 
112+ge::graphStatus FindCustomShapeInferOpFunc(KernelContext *context) {
113+ const char *node_type = context->GetInputValue<char *>(0);
114+ GE_ASSERT_NOTNULL(node_type, "Failed to find custom shape infer op func, node type is nullptr");
115+ auto custom_op_registry = context->GetInputValue<ge::CustomOpRegistry *>(1);
116+ GE_ASSERT_NOTNULL(custom_op_registry, "Failed to find custom shape infer op func, custom op registry is nullptr.");
117+ ge::BaseCustomOp *custom_op_ptr =
118+ custom_op_registry->GetCustomOpCommonCapability(node_type, ge::CustomOpCapability::kShapeInfer);
119+ GE_ASSERT_NOTNULL(custom_op_ptr, "Failed to find custom shape infer op func for op type %s in custom op registry.",
120+ node_type);
121+ auto *shape_infer_op_ptr = ge::CustomOpCast<ge::ShapeInferOp>(custom_op_ptr);
122+ GE_ASSERT_NOTNULL(shape_infer_op_ptr, "Failed to cast custom op %s to ShapeInferOp.", node_type);
123+ auto chain = context->GetOutput(0);
124+ GE_ASSERT_NOTNULL(chain);
125+ chain->Set(shape_infer_op_ptr, nullptr);
126+ return ge::GRAPH_SUCCESS;
127+}
128+ 
93static ge::graphStatus CreateOutputTensors(const ExtendedKernelContext *extended_kernel_context,129static ge::graphStatus CreateOutputTensors(const ExtendedKernelContext *extended_kernel_context,
94 KernelContext *context) {130 KernelContext *context) {
95 const size_t node_output_num = extended_kernel_context->GetComputeNodeOutputNum();131 const size_t node_output_num = extended_kernel_context->GetComputeNodeOutputNum();
@@ -136,9 +172,8 @@ static ge::graphStatus CreateCustomOpOutputs(const ge::FastNode *node, KernelCon
136 return ge::GRAPH_SUCCESS;172 return ge::GRAPH_SUCCESS;
137}173}
138 174 
139-static ge::graphStatus CopyShapeFromTemplateTensors(KernelContext *context, size_t node_input_num,175+static ge::graphStatus CopyShapeFromTemplateTensors(KernelContext *context, size_t template_tensor_start,
140 size_t node_output_num) {176 size_t node_output_num) {
141- const size_t template_tensor_start = node_input_num + static_cast<size_t>(CustomOpInput::kEnd);
142 for (size_t index = 0; index < node_output_num; ++index) {177 for (size_t index = 0; index < node_output_num; ++index) {
143 auto template_tensor = context->GetInputPointer<Tensor>(template_tensor_start + index);178 auto template_tensor = context->GetInputPointer<Tensor>(template_tensor_start + index);
144 auto output_tensor = context->GetOutputPointer<Tensor>(index);179 auto output_tensor = context->GetOutputPointer<Tensor>(index);
@@ -188,11 +223,50 @@ ge::graphStatus ExecuteCustomOpFunc(KernelContext *context) {
188ge::graphStatus ExecuteCustomOpWithInferShapeFunc(KernelContext *context) {223ge::graphStatus ExecuteCustomOpWithInferShapeFunc(KernelContext *context) {
189 auto *eager_context = reinterpret_cast<EagerOpExecutionContext *>(context);224 auto *eager_context = reinterpret_cast<EagerOpExecutionContext *>(context);
190 GE_ASSERT_NOTNULL(eager_context);225 GE_ASSERT_NOTNULL(eager_context);
191- GE_ASSERT_SUCCESS(CopyShapeFromTemplateTensors(context, eager_context->GetComputeNodeInputNum(),226+ const size_t template_tensor_start =
192- eager_context->GetComputeNodeOutputNum()));227+ eager_context->GetComputeNodeInputNum() + static_cast<size_t>(CustomOpInput::kEnd);
228+ GE_ASSERT_SUCCESS(
229+ CopyShapeFromTemplateTensors(context, template_tensor_start, eager_context->GetComputeNodeOutputNum()));
193 return ExecuteCustomOpImpl(context);230 return ExecuteCustomOpImpl(context);
194}231}
195 232 
233+static ge::graphStatus CreateHostCustomOpOutputs(const ge::FastNode *node, KernelContext *context) {
234+ (void)node;
235+ auto *extended_kernel_context = reinterpret_cast<ExtendedKernelContext *>(context);
236+ GE_ASSERT_NOTNULL(extended_kernel_context);
237+ return CreateOutputTensors(extended_kernel_context, context);
238+}
239+ 
240+static ge::graphStatus ExecuteHostCustomOpImpl(KernelContext *context) {
241+ auto *host_context = reinterpret_cast<HostCpuOpExecutionContext *>(context);
242+ GE_ASSERT_NOTNULL(host_context);
243+ const size_t node_input_num = host_context->GetComputeNodeInputNum();
244+ auto custom_op_ptr =
245+ context->GetInputValue<ge::BaseCustomOp *>(node_input_num + static_cast<size_t>(HostCustomOpInput::kFunc));
246+ GE_ASSERT_NOTNULL(custom_op_ptr);
247+ auto *host_execute_op_ptr = ge::CustomOpCast<ge::HostCpuExecuteOp>(custom_op_ptr);
248+ if (host_execute_op_ptr == nullptr) {
249+ GELOGE(ge::FAILED, "%s is host CPU custom op but did not implement HostCpuExecuteOp", host_context->GetNodeType());
250+ return ge::GRAPH_FAILED;
251+ }
252+ GE_ASSERT_SUCCESS(host_execute_op_ptr->Execute(host_context));
253+ return ge::GRAPH_SUCCESS;
254+}
255+ 
256+ge::graphStatus ExecuteHostCustomOpFunc(KernelContext *context) {
257+ return ExecuteHostCustomOpImpl(context);
258+}
259+ 
260+ge::graphStatus ExecuteHostCustomOpWithInferShapeFunc(KernelContext *context) {
261+ auto *host_context = reinterpret_cast<HostCpuOpExecutionContext *>(context);
262+ GE_ASSERT_NOTNULL(host_context);
263+ const size_t template_tensor_start =
264+ host_context->GetComputeNodeInputNum() + static_cast<size_t>(HostCustomOpInput::kEnd);
265+ GE_ASSERT_SUCCESS(
266+ CopyShapeFromTemplateTensors(context, template_tensor_start, host_context->GetComputeNodeOutputNum()));
267+ return ExecuteHostCustomOpImpl(context);
268+}
269+ 
196ge::graphStatus FreeCustomOpWorkspacesFunc(KernelContext *context) {270ge::graphStatus FreeCustomOpWorkspacesFunc(KernelContext *context) {
197 auto memory_vec = context->MutableInputPointer<std::vector<GertMemBlock *>>(0);271 auto memory_vec = context->MutableInputPointer<std::vector<GertMemBlock *>>(0);
198 GE_ASSERT_NOTNULL(memory_vec);272 GE_ASSERT_NOTNULL(memory_vec);
@@ -216,17 +290,18 @@ ge::graphStatus FreeArgsGuarderFunc(KernelContext *context) {
216 return ge::GRAPH_SUCCESS;290 return ge::GRAPH_SUCCESS;
217}291}
218 292 
219-static std::vector<std::string> CustomOpExecuteKernelTrace(const KernelContext *context) {293+template <typename OpContextT>
294+std::vector<std::string> CustomOpExecuteKernelTraceImpl(const KernelContext *context) {
220 auto extend_context = reinterpret_cast<const ExtendedKernelContext *>(context);295 auto extend_context = reinterpret_cast<const ExtendedKernelContext *>(context);
221 auto compute_node_info = extend_context->GetComputeNodeInfo();296 auto compute_node_info = extend_context->GetComputeNodeInfo();
222 if (compute_node_info == nullptr) {297 if (compute_node_info == nullptr) {
223 return {PrintNodeType(context), "compute_node_info is nullptr"};298 return {PrintNodeType(context), "compute_node_info is nullptr"};
224 }299 }
225- auto *eager_op_context = reinterpret_cast<const EagerOpExecutionContext *>(context);300+ auto *op_context = reinterpret_cast<const OpContextT *>(context);
226 std::stringstream input_tensor_ss;301 std::stringstream input_tensor_ss;
227 input_tensor_ss << "input tensor: ";302 input_tensor_ss << "input tensor: ";
228 for (size_t i = 0U; i < compute_node_info->GetInputsNum(); ++i) {303 for (size_t i = 0U; i < compute_node_info->GetInputsNum(); ++i) {
229- auto tensor = eager_op_context->GetInputTensor(i);304+ auto tensor = op_context->GetInputTensor(i);
230 if (tensor == nullptr) {305 if (tensor == nullptr) {
231 return {PrintNodeType(context), "The " + std::to_string(i) + "th's input tensor is nullptr"};306 return {PrintNodeType(context), "The " + std::to_string(i) + "th's input tensor is nullptr"};
232 }307 }
@@ -240,7 +315,7 @@ static std::vector<std::string> CustomOpExecuteKernelTrace(const KernelContext *
240 std::stringstream output_tensor_ss;315 std::stringstream output_tensor_ss;
241 output_tensor_ss << "output tensor: ";316 output_tensor_ss << "output tensor: ";
242 for (size_t i = 0U; i < compute_node_info->GetOutputsNum(); ++i) {317 for (size_t i = 0U; i < compute_node_info->GetOutputsNum(); ++i) {
243- auto tensor = eager_op_context->GetOutputTensor(i);318+ auto tensor = op_context->GetOutputTensor(i);
244 if (tensor == nullptr) {319 if (tensor == nullptr) {
245 return {PrintNodeType(context), "The " + std::to_string(i) + "th's output tensor is nullptr"};320 return {PrintNodeType(context), "The " + std::to_string(i) + "th's output tensor is nullptr"};
246 }321 }
@@ -254,17 +329,18 @@ static std::vector<std::string> CustomOpExecuteKernelTrace(const KernelContext *
254 return {PrintNodeType(context), input_tensor_ss.str(), output_tensor_ss.str(), PrintStreamIdAndTaskId()};329 return {PrintNodeType(context), input_tensor_ss.str(), output_tensor_ss.str(), PrintStreamIdAndTaskId()};
255}330}
256 331 
257-ge::graphStatus CustomOpProfilingDataFill(const KernelContext *context, ProfilingInfoWrapper &prof_info) {332+template <typename OpContextT>
333+ge::graphStatus CustomOpProfilingDataFillImpl(const KernelContext *context, ProfilingInfoWrapper &prof_info) {
258 prof_info.SetBlockDim(std::numeric_limits<uint32_t>::max());334 prof_info.SetBlockDim(std::numeric_limits<uint32_t>::max());
259 auto extend_context = reinterpret_cast<const ExtendedKernelContext *>(context);335 auto extend_context = reinterpret_cast<const ExtendedKernelContext *>(context);
260 auto compute_node_info = extend_context->GetComputeNodeInfo();336 auto compute_node_info = extend_context->GetComputeNodeInfo();
261 GE_ASSERT_NOTNULL(compute_node_info);337 GE_ASSERT_NOTNULL(compute_node_info);
262 auto node_input_num = compute_node_info->GetInputsNum();338 auto node_input_num = compute_node_info->GetInputsNum();
263- const auto eager_context = reinterpret_cast<const EagerOpExecutionContext *>(context);339+ const auto *op_context = reinterpret_cast<const OpContextT *>(context);
264- GE_ASSERT_NOTNULL(eager_context);340+ GE_ASSERT_NOTNULL(op_context);
265 std::vector<std::vector<int64_t>> input_shapes;341 std::vector<std::vector<int64_t>> input_shapes;
266 for (size_t i = 0UL; i < node_input_num; i++) {342 for (size_t i = 0UL; i < node_input_num; i++) {
267- auto tensor = eager_context->GetInputTensor(i);343+ auto tensor = op_context->GetInputTensor(i);
268 GE_ASSERT_NOTNULL(tensor);344 GE_ASSERT_NOTNULL(tensor);
269 auto shape = tensor->GetStorageShape();345 auto shape = tensor->GetStorageShape();
270 std::vector<int64_t> dims;346 std::vector<int64_t> dims;
@@ -276,7 +352,7 @@ ge::graphStatus CustomOpProfilingDataFill(const KernelContext *context, Profilin
276 auto node_output_num = compute_node_info->GetOutputsNum();352 auto node_output_num = compute_node_info->GetOutputsNum();
277 std::vector<std::vector<int64_t>> output_shapes;353 std::vector<std::vector<int64_t>> output_shapes;
278 for (size_t i = 0UL; i < node_output_num; i++) {354 for (size_t i = 0UL; i < node_output_num; i++) {
279- auto tensor = eager_context->GetOutputTensor(i);355+ auto tensor = op_context->GetOutputTensor(i);
280 GE_ASSERT_NOTNULL(tensor);356 GE_ASSERT_NOTNULL(tensor);
281 auto shape = tensor->GetStorageShape();357 auto shape = tensor->GetStorageShape();
282 std::vector<int64_t> dims;358 std::vector<int64_t> dims;
@@ -289,6 +365,22 @@ ge::graphStatus CustomOpProfilingDataFill(const KernelContext *context, Profilin
289 return ge::GRAPH_SUCCESS;365 return ge::GRAPH_SUCCESS;
290}366}
291 367 
368+static std::vector<std::string> CustomOpExecuteKernelTrace(const KernelContext *context) {
369+ return CustomOpExecuteKernelTraceImpl<EagerOpExecutionContext>(context);
370+}
371+ 
372+static std::vector<std::string> HostCustomOpExecuteKernelTrace(const KernelContext *context) {
373+ return CustomOpExecuteKernelTraceImpl<HostCpuOpExecutionContext>(context);
374+}
375+ 
376+static ge::graphStatus CustomOpProfilingDataFill(const KernelContext *context, ProfilingInfoWrapper &prof_info) {
377+ return CustomOpProfilingDataFillImpl<EagerOpExecutionContext>(context, prof_info);
378+}
379+ 
380+static ge::graphStatus HostCustomOpProfilingDataFill(const KernelContext *context, ProfilingInfoWrapper &prof_info) {
381+ return CustomOpProfilingDataFillImpl<HostCpuOpExecutionContext>(context, prof_info);
382+}
383+ 
292REGISTER_KERNEL(FindCustomOp).RunFunc(FindCustomOpFunc);384REGISTER_KERNEL(FindCustomOp).RunFunc(FindCustomOpFunc);
293REGISTER_KERNEL(ExecuteCustomOp)385REGISTER_KERNEL(ExecuteCustomOp)
294 .OutputsCreator(CreateCustomOpOutputs)386 .OutputsCreator(CreateCustomOpOutputs)
@@ -306,5 +398,20 @@ REGISTER_KERNEL(FreeCustomOpWorkspaces)
306 .RunFunc(FreeCustomOpWorkspacesFunc)398 .RunFunc(FreeCustomOpWorkspacesFunc)
307 .ConcurrentCriticalSectionKey(kKernelUseMemory);399 .ConcurrentCriticalSectionKey(kKernelUseMemory);
308REGISTER_KERNEL(FreeArgsGuarder).RunFunc(FreeArgsGuarderFunc).ConcurrentCriticalSectionKey(kKernelUseMemory);400REGISTER_KERNEL(FreeArgsGuarder).RunFunc(FreeArgsGuarderFunc).ConcurrentCriticalSectionKey(kKernelUseMemory);
401+ 
402+REGISTER_KERNEL(FindHostCpuCustomOp).RunFunc(FindHostCpuCustomOpFunc);
403+REGISTER_KERNEL(FindCustomShapeInferOp).RunFunc(FindCustomShapeInferOpFunc);
404+REGISTER_KERNEL(ExecuteHostCustomOp)
405+ .OutputsCreator(CreateHostCustomOpOutputs)
406+ .RunFunc(ExecuteHostCustomOpFunc)
407+ .TracePrinter(HostCustomOpExecuteKernelTrace)
408+ .ProfilingInfoFiller(HostCustomOpProfilingDataFill)
409+ .ConcurrentCriticalSectionKey(kKernelUseMemory);
410+REGISTER_KERNEL(ExecuteHostCustomOpWithInferShape)
411+ .OutputsCreator(CreateHostCustomOpOutputs)
412+ .RunFunc(ExecuteHostCustomOpWithInferShapeFunc)
413+ .TracePrinter(HostCustomOpExecuteKernelTrace)
414+ .ProfilingInfoFiller(HostCustomOpProfilingDataFill)
415+ .ConcurrentCriticalSectionKey(kKernelUseMemory);
309} // namespace kernel416} // namespace kernel
310} // namespace gert417} // namespace gert
@@ -19,6 +19,10 @@ namespace kernel {
19ge::graphStatus FindCustomOpFunc(KernelContext *context);19ge::graphStatus FindCustomOpFunc(KernelContext *context);
20ge::graphStatus ExecuteCustomOpFunc(KernelContext *context);20ge::graphStatus ExecuteCustomOpFunc(KernelContext *context);
21ge::graphStatus FreeCustomOpWorkspacesFunc(KernelContext *context);21ge::graphStatus FreeCustomOpWorkspacesFunc(KernelContext *context);
22+ge::graphStatus FindHostCpuCustomOpFunc(KernelContext *context);
23+ge::graphStatus FindCustomShapeInferOpFunc(KernelContext *context);
24+ge::graphStatus ExecuteHostCustomOpFunc(KernelContext *context);
25+ge::graphStatus ExecuteHostCustomOpWithInferShapeFunc(KernelContext *context);
22} // namespace kernel26} // namespace kernel
23} // namespace gert27} // namespace gert
24 28 
@@ -132,26 +132,26 @@ std::vector<ValueHolderPtr> BuildInferShapeGraph(const ge::NodePtr &node,
132 return ValueHolder::CreateDataOutput("InferShape", inputs, node->GetAllOutDataAnchorsSize());132 return ValueHolder::CreateDataOutput("InferShape", inputs, node->GetAllOutDataAnchorsSize());
133}133}
134 134 
135-bg::ValueHolderPtr FindCustomOpFunc(const ge::NodePtr &node, LoweringGlobalData &global_data) {135+bg::ValueHolderPtr FindCustomShapeInferOpFunc(const ge::NodePtr &node, LoweringGlobalData &global_data) {
136 auto builder = [&node, &global_data]() -> std::vector<bg::ValueHolderPtr> {136 auto builder = [&node, &global_data]() -> std::vector<bg::ValueHolderPtr> {
137 return bg::FrameSelector::OnInitRoot([&node, &global_data]() -> std::vector<bg::ValueHolderPtr> {137 return bg::FrameSelector::OnInitRoot([&node, &global_data]() -> std::vector<bg::ValueHolderPtr> {
138 auto node_type = ValueHolder::CreateConst(node->GetTypePtr(), node->GetType().size() + 1, true);138 auto node_type = ValueHolder::CreateConst(node->GetTypePtr(), node->GetType().size() + 1, true);
139 ge::CustomOpRegistry *custom_op_registry = global_data.GetCustomOpRegistry().get();139 ge::CustomOpRegistry *custom_op_registry = global_data.GetCustomOpRegistry().get();
140 auto registry_holder = ValueHolder::CreateConst(&custom_op_registry, sizeof(ge::CustomOpRegistry *));140 auto registry_holder = ValueHolder::CreateConst(&custom_op_registry, sizeof(ge::CustomOpRegistry *));
141- return {ValueHolder::CreateSingleDataOutput("FindCustomOp", {node_type, registry_holder})};141+ return {ValueHolder::CreateSingleDataOutput("FindCustomShapeInferOp", {node_type, registry_holder})};
142 });142 });
143 };143 };
144- return global_data.GetOrCreateUniqueValueHolder(node->GetType() + "_FindCustomOp_", builder)[0];144+ return global_data.GetOrCreateUniqueValueHolder(node->GetType() + "_FindCustomShapeInferOp_", builder)[0];
145}145}
146 146 
147std::vector<ValueHolderPtr> BuildCustomOpInferShapeGraph(const ge::NodePtr &node,147std::vector<ValueHolderPtr> BuildCustomOpInferShapeGraph(const ge::NodePtr &node,
148 const std::vector<ValueHolderPtr> &input_shapes,148 const std::vector<ValueHolderPtr> &input_shapes,
149 LoweringGlobalData &global_data) {149 LoweringGlobalData &global_data) {
150- auto custom_op_func = FindCustomOpFunc(node, global_data);150+ auto shape_infer_op_func = FindCustomShapeInferOpFunc(node, global_data);
151 auto infer_shape_func = kernel::InferCustomOpShapeFromInput;151 auto infer_shape_func = kernel::InferCustomOpShapeFromInput;
152 auto infer_shape_func_holder = ValueHolder::CreateConst(&infer_shape_func, sizeof(decltype(infer_shape_func)));152 auto infer_shape_func_holder = ValueHolder::CreateConst(&infer_shape_func, sizeof(decltype(infer_shape_func)));
153 auto inputs = input_shapes;153 auto inputs = input_shapes;
154- inputs.emplace_back(custom_op_func);154+ inputs.emplace_back(shape_infer_op_func);
155 inputs.emplace_back(infer_shape_func_holder);155 inputs.emplace_back(infer_shape_func_holder);
156 return ValueHolder::CreateDataOutput("InferShape", inputs, node->GetAllOutDataAnchorsSize());156 return ValueHolder::CreateDataOutput("InferShape", inputs, node->GetAllOutDataAnchorsSize());
157}157}
@@ -74,9 +74,7 @@ inline ge::graphStatus InferCustomOpShapeFromInput(InferShapeContext *context) {
74 GE_ASSERT_NOTNULL(kernel_context);74 GE_ASSERT_NOTNULL(kernel_context);
75 const auto input_num = kernel_context->GetInputNum();75 const auto input_num = kernel_context->GetInputNum();
76 GE_ASSERT(input_num > 1U);76 GE_ASSERT(input_num > 1U);
77- auto custom_op = kernel_context->GetInputValue<ge::BaseCustomOp *>(input_num - 2U);77+ auto shape_infer_op = kernel_context->GetInputValue<ge::ShapeInferOp *>(input_num - 2U);
78- GE_ASSERT_NOTNULL(custom_op);
79- auto shape_infer_op = ge::CustomOpCast<ge::ShapeInferOp>(custom_op);
80 if (shape_infer_op == nullptr) {78 if (shape_infer_op == nullptr) {
81 GELOGE(ge::GRAPH_FAILED, "Custom op does not implement ShapeInferOp.");79 GELOGE(ge::GRAPH_FAILED, "Custom op does not implement ShapeInferOp.");
82 return ge::GRAPH_FAILED;80 return ge::GRAPH_FAILED;
@@ -36,13 +36,49 @@
36#include "lowering/placement/placed_lowering_register.h"36#include "lowering/placement/placed_lowering_register.h"
37#include "framework/common/ge_types.h"37#include "framework/common/ge_types.h"
38#include "runtime/gert_api.h"38#include "runtime/gert_api.h"
39+#include "exe_graph/runtime/host_cpu_op_execution_context.h"
40+#include "exe_graph/lowering/kernel_run_context_builder.h"
39#include "check/executor_statistician.h"41#include "check/executor_statistician.h"
40#include "faker/nodes_faker_for_exe.h"42#include "faker/nodes_faker_for_exe.h"
41#include "register/op_impl_registry.h"43#include "register/op_impl_registry.h"
42#include "engine/aicpu/graph_builder/bg_aicpu_arg.h"44#include "engine/aicpu/graph_builder/bg_aicpu_arg.h"
45+#include "engine/custom/converter/custom_node_converter.h"
46+#include "engine/gelocal/inputs_converter.h"
47+#include "graph/custom_op_factory.h"
48+#include "graph/custom_op.h"
49+#include "graph/debug/ge_attr_define.h"
50+#include "ge/ut/ge/runtime/fast_v2/common/const_data_helper.h"
43 51 
44using namespace ge;52using namespace ge;
45namespace gert {53namespace gert {
54+class StHostCpuLoweringOp final : public ge::HostCpuExecuteOp {
55+ public:
56+ ge::graphStatus Execute(gert::HostCpuOpExecutionContext *ctx) override {
57+ (void)ctx;
58+ return ge::GRAPH_SUCCESS;
59+ }
60+};
61+ 
62+class StHostCpuShapeInferLoweringOp final : public ge::HostCpuExecuteOp, public ge::ShapeInferOp {
63+ public:
64+ ge::graphStatus Execute(gert::HostCpuOpExecutionContext *) override {
65+ return ge::GRAPH_SUCCESS;
66+ }
67+ 
68+ ge::graphStatus InferShape(gert::InferShapeContext *context) override {
69+ auto input = context->GetInputShape(0);
70+ auto output = context->GetOutputShape(0);
71+ GE_ASSERT_NOTNULL(input);
72+ GE_ASSERT_NOTNULL(output);
73+ *output = *input;
74+ return ge::GRAPH_SUCCESS;
75+ }
76+ 
77+ ge::graphStatus InferDataType(gert::InferDataTypeContext *) override {
78+ return ge::GRAPH_SUCCESS;
79+ }
80+};
81+ 
46class PlaceLoweringResultSystemTest : public bg::BgTest {82class PlaceLoweringResultSystemTest : public bg::BgTest {
47 protected:83 protected:
48 void SetUp() override {84 void SetUp() override {
@@ -103,6 +139,20 @@ LowerResult FakeConverterForCast(const ge::NodePtr &node, const LowerInput &lowe
103 return FakedDeviceConverterWithOrderedHoldersWithType(node, lower_input, "LaunchCast");139 return FakedDeviceConverterWithOrderedHoldersWithType(node, lower_input, "LaunchCast");
104}140}
105 141 
142+void PrepareCustomOpInputs(const ge::ComputeGraphPtr &graph, LoweringGlobalData &global_data,
143+ const std::vector<std::string> &data_names, std::vector<bg::ValueHolderPtr> &shapes,
144+ std::vector<bg::DevMemValueHolderPtr> &addrs) {
145+ bg::LowerConstDataNode(global_data);
146+ for (const auto &name : data_names) {
147+ auto data_ret = LoweringDataNode(graph->FindNode(name), {{}, {}, &global_data});
148+ ASSERT_TRUE(data_ret.result.IsSuccess());
149+ shapes.emplace_back(data_ret.out_shapes[0]);
150+ addrs.emplace_back(data_ret.out_addrs[0]);
151+ graph->FindNode(name)->GetOpDesc()->SetExtAttr("_lowering_result",
152+ PlacedLoweringResult(graph->FindNode(name), std::move(data_ret)));
153+ }
154+}
155+ 
106REG_OP(Add)156REG_OP(Add)
107 .INPUT(x1, TensorType({DT_FLOAT, DT_INT32, DT_INT64, DT_FLOAT16, DT_INT16, DT_INT8, DT_UINT8, DT_DOUBLE,157 .INPUT(x1, TensorType({DT_FLOAT, DT_INT32, DT_INT64, DT_FLOAT16, DT_INT16, DT_INT8, DT_UINT8, DT_DOUBLE,
108 DT_COMPLEX128, DT_COMPLEX64, DT_STRING}))158 DT_COMPLEX128, DT_COMPLEX64, DT_STRING}))
@@ -113,6 +163,136 @@ REG_OP(Add)
113 .OP_END_FACTORY_REG(Add)163 .OP_END_FACTORY_REG(Add)
114} // namespace164} // namespace
115 165 
166+TEST_F(PlaceLoweringResultStringTest, HostCpuCustomNodeLoweringBuildsHostTensors) {
167+ const AscendString op_type("StHostCpuLoweringOp");
168+ CustomOpFactory::RemoveCustomOps({op_type});
169+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
170+ op_type, OpBackend::kHostCPU,
171+ []() -> std::unique_ptr<ge::BaseCustomOp> { return std::make_unique<StHostCpuLoweringOp>(); }),
172+ ge::GRAPH_SUCCESS);
173+ 
174+ auto graph = ShareGraph::BuildCustomOpGraph();
175+ auto custom_op = graph->FindNode("custom_op");
176+ ASSERT_NE(custom_op, nullptr);
177+ custom_op->GetOpDesc()->SetType(op_type.GetString());
178+ auto root_model = GeModelBuilder(graph).BuildGeRootModel();
179+ auto global_data = GlobalDataFaker(root_model).Build();
180+ global_data.SetCustomOpRegistry(CustomOpFactory::GetGlobalRegistryPtr());
181+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kInit);
182+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kMain);
183+ std::vector<bg::ValueHolderPtr> shapes;
184+ std::vector<bg::DevMemValueHolderPtr> addrs;
185+ PrepareCustomOpInputs(graph, global_data, {"data0", "data1", "data2"}, shapes, addrs);
186+ 
187+ auto ret = LoweringHostCustomNode(custom_op, {shapes, addrs, &global_data});
188+ ASSERT_TRUE(ret.result.IsSuccess());
189+ ASSERT_EQ(ret.out_addrs.size(), 1U);
190+ EXPECT_EQ(ret.out_addrs[0]->GetPlacement(), kOnHost);
191+ EXPECT_NE(ExecuteGraphUtils::FindFirstNodeMatchType(bg::ValueHolder::GetCurrentFrame()->GetExecuteGraph().get(),
192+ "ExecuteHostCustomOp"),
193+ nullptr);
194+ CustomOpFactory::RemoveCustomOps({op_type});
195+}
196+ 
197+TEST_F(PlaceLoweringResultStringTest, HostCpuCustomNodeLoweringSkipsUnconnectedInput) {
198+ const AscendString op_type("StHostCpuOptionalLoweringOp");
199+ CustomOpFactory::RemoveCustomOps({op_type});
200+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
201+ op_type, OpBackend::kHostCPU,
202+ []() -> std::unique_ptr<ge::BaseCustomOp> { return std::make_unique<StHostCpuLoweringOp>(); }),
203+ ge::GRAPH_SUCCESS);
204+ 
205+ auto graph = ShareGraph::BuildCustomOpGraph();
206+ auto custom_op = graph->FindNode("custom_op");
207+ ASSERT_NE(custom_op, nullptr);
208+ custom_op->GetOpDesc()->SetType(op_type.GetString());
209+ ASSERT_EQ(ge::GraphUtils::RemoveEdge(graph->FindNode("data2")->GetOutDataAnchor(0), custom_op->GetInDataAnchor(2)),
210+ ge::GRAPH_SUCCESS);
211+ auto root_model = GeModelBuilder(graph).BuildGeRootModel();
212+ auto global_data = GlobalDataFaker(root_model).Build();
213+ global_data.SetCustomOpRegistry(CustomOpFactory::GetGlobalRegistryPtr());
214+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kInit);
215+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kMain);
216+ std::vector<bg::ValueHolderPtr> shapes;
217+ std::vector<bg::DevMemValueHolderPtr> addrs;
218+ PrepareCustomOpInputs(graph, global_data, {"data0", "data1"}, shapes, addrs);
219+ auto ret = LoweringHostCustomNode(custom_op, {shapes, addrs, &global_data});
220+ ASSERT_TRUE(ret.result.IsSuccess());
221+ ASSERT_EQ(ret.out_addrs.size(), 1U);
222+ CustomOpFactory::RemoveCustomOps({op_type});
223+}
224+ 
225+TEST_F(PlaceLoweringResultStringTest, HostCpuCustomNodeLoweringAddsInputGuardDependency) {
226+ const AscendString op_type("StHostCpuGuardLoweringOp");
227+ CustomOpFactory::RemoveCustomOps({op_type});
228+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
229+ op_type, OpBackend::kHostCPU,
230+ []() -> std::unique_ptr<ge::BaseCustomOp> { return std::make_unique<StHostCpuLoweringOp>(); }),
231+ ge::GRAPH_SUCCESS);
232+ auto graph = ShareGraph::BuildCustomOpGraph();
233+ auto custom_op = graph->FindNode("custom_op");
234+ custom_op->GetOpDesc()->SetType(op_type.GetString());
235+ auto root_model = GeModelBuilder(graph).BuildGeRootModel();
236+ auto global_data = GlobalDataFaker(root_model).Build();
237+ global_data.SetCustomOpRegistry(CustomOpFactory::GetGlobalRegistryPtr());
238+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kInit);
239+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kMain);
240+ std::vector<bg::ValueHolderPtr> shapes;
241+ std::vector<bg::DevMemValueHolderPtr> addrs;
242+ PrepareCustomOpInputs(graph, global_data, {"data0", "data1", "data2"}, shapes, addrs);
243+ auto input_guarder = bg::ValueHolder::CreateVoidGuarder("FreeInput", addrs[0], {});
244+ ASSERT_NE(input_guarder, nullptr);
245+ addrs[0]->SetGuarder(input_guarder);
246+ auto ret = LoweringHostCustomNode(custom_op, {shapes, addrs, &global_data});
247+ ASSERT_TRUE(ret.result.IsSuccess());
248+ CustomOpFactory::RemoveCustomOps({op_type});
249+}
250+ 
251+TEST_F(PlaceLoweringResultStringTest, HostCpuCustomNodeLoweringUsesInferShapeKernel) {
252+ const AscendString op_type("StHostCpuShapeInferLoweringOp");
253+ CustomOpFactory::RemoveCustomOps({op_type});
254+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(op_type, OpBackend::kHostCPU,
255+ []() -> std::unique_ptr<ge::BaseCustomOp> {
256+ return std::make_unique<StHostCpuShapeInferLoweringOp>();
257+ }),
258+ ge::GRAPH_SUCCESS);
259+ 
260+ auto graph = ShareGraph::BuildCustomOpGraph();
261+ auto custom_op = graph->FindNode("custom_op");
262+ ASSERT_NE(custom_op, nullptr);
263+ custom_op->GetOpDesc()->SetType(op_type.GetString());
264+ auto root_model = GeModelBuilder(graph).BuildGeRootModel();
265+ auto global_data = GlobalDataFaker(root_model).Build();
266+ global_data.SetCustomOpRegistry(CustomOpFactory::GetGlobalRegistryPtr());
267+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kInit);
268+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kMain);
269+ std::vector<bg::ValueHolderPtr> shapes;
270+ std::vector<bg::DevMemValueHolderPtr> addrs;
271+ PrepareCustomOpInputs(graph, global_data, {"data0", "data1", "data2"}, shapes, addrs);
272+ auto input_guarder = bg::ValueHolder::CreateVoidGuarder("FreeInput", addrs[0], {});
273+ ASSERT_NE(input_guarder, nullptr);
274+ addrs[0]->SetGuarder(input_guarder);
275+ 
276+ auto ret = LoweringHostCustomNode(custom_op, {shapes, addrs, &global_data});
277+ ASSERT_TRUE(ret.result.IsSuccess());
278+ EXPECT_NE(ExecuteGraphUtils::FindFirstNodeMatchType(bg::ValueHolder::GetCurrentFrame()->GetExecuteGraph().get(),
279+ "ExecuteHostCustomOpWithInferShape"),
280+ nullptr);
281+ CustomOpFactory::RemoveCustomOps({op_type});
282+}
283+ 
284+TEST_F(PlaceLoweringResultStringTest, HostCpuContextMakeOutputRefInputUsesMutableOutput) {
285+ auto op_desc = std::make_shared<ge::OpDesc>("host_ref", "UnknownHostRefOp");
286+ ASSERT_EQ(op_desc->AddInputDesc(ge::GeTensorDesc(ge::GeShape({1}), ge::FORMAT_ND, ge::DT_FLOAT)), ge::GRAPH_SUCCESS);
287+ ASSERT_EQ(op_desc->AddOutputDesc(ge::GeTensorDesc(ge::GeShape({1}), ge::FORMAT_ND, ge::DT_FLOAT)), ge::GRAPH_SUCCESS);
288+ gert::Tensor input;
289+ gert::Tensor output;
290+ auto holder = gert::KernelRunContextBuilder().Inputs({&input, nullptr}).Outputs({&output}).Build(op_desc);
291+ auto *context = reinterpret_cast<gert::HostCpuOpExecutionContext *>(holder.GetKernelContext());
292+ ASSERT_NE(context, nullptr);
293+ EXPECT_EQ(context->MakeOutputRefInput(0, 0), &output);
294+}
295+ 
116TEST_F(PlaceLoweringResultSystemTest, H2DRunAfterLaunch_PlacedLoweringResult) {296TEST_F(PlaceLoweringResultSystemTest, H2DRunAfterLaunch_PlacedLoweringResult) {
117 GertRuntimeStub runtime_stub;297 GertRuntimeStub runtime_stub;
118 std::string host_kernel_lib_faker = "host_cpu_kernel_stub";298 std::string host_kernel_lib_faker = "host_cpu_kernel_stub";
@@ -20,8 +20,8 @@
20#include <string>20#include <string>
21#include <unistd.h>21#include <unistd.h>
22#include "common/share_graph.h"22#include "common/share_graph.h"
23-#include "faker/global_data_faker.h"
24#include "faker/fake_value.h"23#include "faker/fake_value.h"
24+#include "faker/space_registry_faker.h"
25#include "rt_external_base.h"25#include "rt_external_base.h"
26#include "ge/ge_api.h"26#include "ge/ge_api.h"
27#include "ge/ge_api_error_codes.h"27#include "ge/ge_api_error_codes.h"
@@ -53,6 +53,7 @@
53#include "hcom/hcom_topo_info.h"53#include "hcom/hcom_topo_info.h"
54#include "common/opskernel/ops_kernel_info_types.h"54#include "common/opskernel/ops_kernel_info_types.h"
55#include "engines/custom_engine/custom_graph_optimizer.h"55#include "engines/custom_engine/custom_graph_optimizer.h"
56+#include "engines/custom_engine/custom_ops_kernel_info_store.h"
56#include "engines/custom_engine/custom_ops_kernel_builder.h"57#include "engines/custom_engine/custom_ops_kernel_builder.h"
57#include "graph/compute_graph.h"58#include "graph/compute_graph.h"
58#include "graph/custom_op/cast.h"59#include "graph/custom_op/cast.h"
@@ -69,7 +70,6 @@
69#include "runtime/custom_op/custom_op_loader.h"70#include "runtime/custom_op/custom_op_loader.h"
70#include "runtime/custom_op/python_custom_op_bridge_loader.h"71#include "runtime/custom_op/python_custom_op_bridge_loader.h"
71#include "exe_graph/runtime/storage_shape.h"72#include "exe_graph/runtime/storage_shape.h"
72-#include "exe_graph/runtime/gert_mem_allocator.h"
73#include "faker/kernel_run_context_facker.h"73#include "faker/kernel_run_context_facker.h"
74#include "register/kernel_registry.h"74#include "register/kernel_registry.h"
75#include "runtime/v2/kernel/common_kernel_impl/infer_shape.h"75#include "runtime/v2/kernel/common_kernel_impl/infer_shape.h"
@@ -92,6 +92,13 @@ REG_OP(StPythonCompilableCustomOp)
92 .OUTPUT(z, TensorType::ALL())92 .OUTPUT(z, TensorType::ALL())
93 .REQUIRED_ATTR(bias, Int)93 .REQUIRED_ATTR(bias, Int)
94 .OP_END_FACTORY_REG(StPythonCompilableCustomOp);94 .OP_END_FACTORY_REG(StPythonCompilableCustomOp);
95+ 
96+REG_OP(StHostCpuE2ECustomOp)
97+ .INPUT(x0, TensorType::ALL())
98+ .INPUT(x1, TensorType::ALL())
99+ .INPUT(x2, TensorType::ALL())
100+ .OUTPUT(y, TensorType::ALL())
101+ .OP_END_FACTORY_REG(StHostCpuE2ECustomOp);
95} // namespace ge102} // namespace ge
96 103 
97namespace ge {104namespace ge {
@@ -712,6 +719,7 @@ class CustomOpRefreshTest : public testing::Test {
712 void TearDown() {719 void TearDown() {
713 OpsKernelBuilderRegistry::GetInstance().Unregister("AiCoreLib");720 OpsKernelBuilderRegistry::GetInstance().Unregister("AiCoreLib");
714 OpsKernelBuilderRegistry::GetInstance().Unregister("RTSLib");721 OpsKernelBuilderRegistry::GetInstance().Unregister("RTSLib");
722+ TearDownForGenerateTask(kCustomOpKernelLibName);
715 }723 }
716};724};
717 725 
@@ -774,47 +782,6 @@ class InferMetaCoverageCustomOpForSt final : public CustomOpInferMetaProvider {
774 }782 }
775};783};
776 784 
777-class HostCpuStAllocator final : public gert::GertAllocator {
778- public:
779- HostCpuStAllocator() : GertAllocator(-1, gert::kOnHost) {}
780- 
781- gert::GertMemBlock *Malloc(size_t) override {
782- return nullptr;
783- }
784- 
785- gert::GertTensorData MallocTensorData(size_t) override {
786- return {};
787- }
788- 
789- gert::TensorData MallocTensorDataFromL1(size_t size) override {
790- std::unique_ptr<uint8_t[]> block(new uint8_t[size]);
791- auto *address = block.get();
792- blocks_.emplace_back(std::move(block));
793- return gert::TensorData(address, nullptr, size, gert::kOnHost);
794- }
795- 
796- void Free(gert::GertMemBlock *) override {}
797- 
798- ge::graphStatus FreeAt(int64_t, gert::GertMemBlock *) override {
799- return ge::GRAPH_SUCCESS;
800- }
801- 
802- ge::graphStatus ShareFromTensorData(const gert::TensorData &, gert::GertTensorData &) override {
803- return ge::GRAPH_SUCCESS;
804- }
805- 
806- int64_t GetStreamNum() override {
807- return 0;
808- }
809- 
810- ge::graphStatus SetL1Allocator(ge::Allocator *) override {
811- return ge::GRAPH_SUCCESS;
812- }
813- 
814- private:
815- std::vector<std::unique_ptr<uint8_t[]>> blocks_;
816-};
817- 
818class StRegistryShapeInferOp : public ShapeInferOp {785class StRegistryShapeInferOp : public ShapeInferOp {
819 public:786 public:
820 graphStatus InferShape(gert::InferShapeContext *) override {787 graphStatus InferShape(gert::InferShapeContext *) override {
@@ -828,6 +795,23 @@ class StRegistryShapeInferOp : public ShapeInferOp {
828 795 
829class StRegistryShapeInferOpOther final : public StRegistryShapeInferOp {};796class StRegistryShapeInferOpOther final : public StRegistryShapeInferOp {};
830 797 
798+class StHostCpuE2ECustomOp final : public HostCpuExecuteOp {
799+ public:
800+ graphStatus Execute(gert::HostCpuOpExecutionContext *ctx) override {
801+ execute_called_ = true;
802+ auto *input = ctx->GetInputTensor(0);
803+ if (input == nullptr) {
804+ return GRAPH_FAILED;
805+ }
806+ auto *output = ctx->MallocOutputTensor(0, input->GetShape(), input->GetFormat(), input->GetDataType());
807+ return (output == nullptr || output->GetPlacement() != gert::kOnHost) ? GRAPH_FAILED : GRAPH_SUCCESS;
808+ }
809+ 
810+ static bool execute_called_;
811+};
812+ 
813+bool StHostCpuE2ECustomOp::execute_called_ = false;
814+ 
831class TestBaseCustomOp : public EagerExecuteOp {815class TestBaseCustomOp : public EagerExecuteOp {
832 public:816 public:
833 graphStatus Execute(gert::EagerOpExecutionContext *ctx) override {817 graphStatus Execute(gert::EagerOpExecutionContext *ctx) override {
@@ -2097,6 +2081,8 @@ TEST_F(CustomOpFactoryStTest, PythonCustomOpInferMetaRunsThroughRt2WithNativeAtt
2097 auto *base_op =2081 auto *base_op =
2098 CustomOpFactory::CreateOrGetCustomOp(AscendString(kPythonRt2InferMetaOpTypeForSt), OpBackend::kDevice);2082 CustomOpFactory::CreateOrGetCustomOp(AscendString(kPythonRt2InferMetaOpTypeForSt), OpBackend::kDevice);
2099 ASSERT_NE(base_op, nullptr);2083 ASSERT_NE(base_op, nullptr);
2084+ auto *shape_infer_op = CustomOpCast<ShapeInferOp>(base_op);
2085+ ASSERT_NE(shape_infer_op, nullptr);
2100 2086 
2101 gert::StorageShape input_shape({7, 13}, {7, 13});2087 gert::StorageShape input_shape({7, 13}, {7, 13});
2102 gert::Tensor output;2088 gert::Tensor output;
@@ -2127,7 +2113,7 @@ TEST_F(CustomOpFactoryStTest, PythonCustomOpInferMetaRunsThroughRt2WithNativeAtt
2127 {"attr_list_str", AnyValue::CreateFrom<std::vector<std::string>>({"a", "b"})},2113 {"attr_list_str", AnyValue::CreateFrom<std::vector<std::string>>({"a", "b"})},
2128 {"attr_list_dtype", AnyValue::CreateFrom<std::vector<DataType>>({DT_FLOAT, DT_INT32})},2114 {"attr_list_dtype", AnyValue::CreateFrom<std::vector<DataType>>({DT_FLOAT, DT_INT32})},
2129 {"attr_list_list_int", AnyValue::CreateFrom<std::vector<std::vector<int64_t>>>({{3, 4}, {5}})}})2115 {"attr_list_list_int", AnyValue::CreateFrom<std::vector<std::vector<int64_t>>>({{3, 4}, {5}})}})
2130- .Inputs({&input_shape, base_op, reinterpret_cast<void *>(infer_shape_func)})2116+ .Inputs({&input_shape, shape_infer_op, reinterpret_cast<void *>(infer_shape_func)})
2131 .Outputs({&output})2117 .Outputs({&output})
2132 .Build();2118 .Build();
2133 2119 
@@ -2159,36 +2145,67 @@ TEST_F(CustomOpFactoryStTest, CustomOpInferMetaCompilePath) {
2159 CustomOpFactory::RemoveCustomOps({AscendString(kInferMetaCoverageOpTypeForSt)});2145 CustomOpFactory::RemoveCustomOps({AscendString(kInferMetaCoverageOpTypeForSt)});
2160}2146}
2161 2147 
2162-TEST_F(CustomOpFactoryStTest, HostCpuOpExecutionContextMinimalPaths) {2148+TEST_F(CustomOpRefreshTest, HostCpuCustomOpSessionRun) {
2163- gert::Tensor input_tensor = {2149+ const AscendString op_type("StHostCpuE2ECustomOp");
2164- {{2, 2}, {2, 2}}, {FORMAT_ND, FORMAT_ND, {}}, gert::kOnHost, DT_FLOAT, reinterpret_cast<void *>(0x12345)};2150+ MockForGenerateTask(kCustomOpKernelLibName, GenerateTaskForCustomOp);
2165- gert::Tensor output_tensor;2151+ CustomOpFactory::RemoveCustomOps({op_type});
2166- HostCpuStAllocator allocator;2152+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2167- auto context_holder = gert::KernelRunContextFaker()2153+ op_type, OpBackend::kHostCPU,
2168- .NodeIoNum(1, 1)2154+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StHostCpuE2ECustomOp>(); }),
2169- .IrInputNum(1)2155+ GRAPH_SUCCESS);
2170- .NodeInputTd(0, DT_FLOAT, FORMAT_ND, FORMAT_ND)2156+ auto &store_map = const_cast<std::map<std::string, OpsKernelInfoStorePtr> &>(
2171- .NodeOutputTd(0, DT_FLOAT, FORMAT_ND, FORMAT_ND)2157+ OpsKernelManager::GetInstance().GetAllOpsKernelInfoStores());
2172- .Inputs({&input_tensor, &allocator})2158+ if (store_map.find(kCustomOpKernelLibName) == store_map.end()) {
2173- .Outputs({&output_tensor})2159+ store_map.emplace(kCustomOpKernelLibName, std::make_shared<custom::CustomOpsKernelInfoStore>());
2174- .Build();2160+ }
2175- auto *context = context_holder.GetContext<gert::HostCpuOpExecutionContext>();2161+ ASSERT_NE(store_map.find(kCustomOpKernelLibName), store_map.end());
2162+ ASSERT_EQ(OpsKernelManager::GetInstance().RefreshOpsKernelInfo(), SUCCESS);
2176 2163 
2177- ASSERT_NE(context, nullptr);2164+ setenv("ENABLE_RUNTIME_V2", "1", 1);
2178- EXPECT_EQ(context->GetInputTensor(0), &input_tensor);2165+ auto compute_graph = ShareGraph::BuildOnlyCustomOpKnowShapeGraph();
2179- EXPECT_EQ(context->GetOutputTensor(0), &output_tensor);2166+ ASSERT_NE(compute_graph, nullptr);
2167+ auto custom_node = compute_graph->FindNode("custom_op");
2168+ ASSERT_NE(custom_node, nullptr);
2169+ custom_node->GetOpDesc()->SetType(op_type.GetString());
2170+ custom_node->GetOpDesc()->SetOpKernelLibName("");
2171+ custom_node->GetOpDesc()->SetOpEngineName("");
2172+ auto graph = GraphUtilsEx::CreateGraphFromComputeGraph(compute_graph);
2180 2173 
2181- auto *allocated_output = context->MallocOutputTensor(0, {{2, 2}, {2, 2}}, {FORMAT_ND, FORMAT_ND, {}}, DT_FLOAT);2174+ std::map<AscendString, AscendString> options;
2182- ASSERT_NE(allocated_output, nullptr);2175+ options.emplace(ge::OPTION_CONST_LIFECYCLE, "graph");
2183- EXPECT_EQ(allocated_output->GetPlacement(), gert::kOnHost);2176+ options.emplace(ge::OPTION_GRAPH_RUN_MODE, "0");
2184- EXPECT_EQ(allocated_output->GetSize(), 512U);2177+ Session session(options);
2185- EXPECT_NE(allocated_output->GetAddr(), nullptr);2178+ constexpr uint32_t graph_id = 1U;
2179+ ASSERT_EQ(session.AddGraph(graph_id, graph), SUCCESS);
2186 2180 
2187- auto *ref_output = context->MakeOutputRefInput(0, 0);2181+ std::vector<ge::Tensor> inputs;
2188- ASSERT_NE(ref_output, nullptr);2182+ std::vector<ge::Tensor> outputs;
2189- EXPECT_EQ(ref_output->GetOriginShape(), input_tensor.GetOriginShape());2183+ ConstructCustomInputOutputTensor(3, 1, inputs, outputs);
2190- EXPECT_EQ(ref_output->GetStorageShape(), input_tensor.GetStorageShape());2184+ StHostCpuE2ECustomOp::execute_called_ = false;
2191- EXPECT_EQ(ref_output->GetAddr(), input_tensor.GetAddr());2185+ ASSERT_EQ(session.RunGraph(graph_id, inputs, outputs), SUCCESS);
2186+ EXPECT_TRUE(StHostCpuE2ECustomOp::execute_called_);
2187+ 
2188+ session.RemoveGraph(graph_id);
2189+ unsetenv("ENABLE_RUNTIME_V2");
2190+ CustomOpFactory::RemoveCustomOps({op_type});
2191+}
2192+ 
2193+TEST_F(CustomOpFactoryStTest, FindCustomShapeInferKernelUsesHostRegistry) {
2194+ const AscendString op_type("StShapeInferKernelOp");
2195+ CustomOpFactory::RemoveCustomOps({op_type});
2196+ auto registry = std::make_shared<CustomOpRegistry>();
2197+ ASSERT_EQ(registry->RegisterCreator(
2198+ op_type, OpBackend::kDevice,
2199+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StRegistryShapeInferOp>(); }),
2200+ GRAPH_SUCCESS);
2201+ auto run_context = gert::BuildKernelRunContext(2, 1);
2202+ run_context.value_holder[0].Set(const_cast<char *>(op_type.GetString()), nullptr);
2203+ run_context.value_holder[1].Set(registry.get(), nullptr);
2204+ const auto *funcs = gert::KernelRegistry::GetInstance().FindKernelFuncs("FindCustomShapeInferOp");
2205+ ASSERT_NE(funcs, nullptr);
2206+ EXPECT_EQ(funcs->run_func(run_context), GRAPH_SUCCESS);
2207+ EXPECT_NE(*run_context.GetContext<gert::KernelContext>()->GetOutputPointer<ShapeInferOp *>(0), nullptr);
2208+ CustomOpFactory::RemoveCustomOps({op_type});
2192}2209}
2193 2210 
2194TEST_F(CustomOpFactoryStTest, CustomOpRegistryCoversCompatibilityAndBackendPaths) {2211TEST_F(CustomOpFactoryStTest, CustomOpRegistryCoversCompatibilityAndBackendPaths) {
@@ -25,6 +25,8 @@
25#include "engines/manager/opskernel_manager/ops_kernel_builder_manager.h"25#include "engines/manager/opskernel_manager/ops_kernel_builder_manager.h"
26#include "engines/manager/opskernel_manager/dnn_ops_kernel_manager.h"26#include "engines/manager/opskernel_manager/dnn_ops_kernel_manager.h"
27#include "graph/partition/optimizer/hostcpu_engine_update_pass.h"27#include "graph/partition/optimizer/hostcpu_engine_update_pass.h"
28+#include "graph/partition/engine_partitioner.h"
29+#include "graph/passes/shape_optimize/mark_graph_unknown_status_pass.h"
28#include "hybrid/common/npu_memory_allocator.h"30#include "hybrid/common/npu_memory_allocator.h"
29#include "graph/bin_cache/node_compile_cache_module.h"31#include "graph/bin_cache/node_compile_cache_module.h"
30#include "register/op_tiling_registry.h"32#include "register/op_tiling_registry.h"
@@ -36,6 +38,8 @@
36#include "macro_utils/dt_public_unscope.h"38#include "macro_utils/dt_public_unscope.h"
37 39 
38#include "graph/operator_reg.h"40#include "graph/operator_reg.h"
41+#include "graph/custom_op_factory.h"
42+#include "graph/custom_op.h"
39#include "graph/ge_attr_value.h"43#include "graph/ge_attr_value.h"
40#include "common/dump/dump_manager.h"44#include "common/dump/dump_manager.h"
41#include "register/op_tiling_registry.h"45#include "register/op_tiling_registry.h"
@@ -77,6 +81,8 @@
77#include "depends/profiler/src/profiling_test_util.h"81#include "depends/profiler/src/profiling_test_util.h"
78#include "graph/manager/host_mem_manager.h"82#include "graph/manager/host_mem_manager.h"
79#include "register/register_custom_pass.h"83#include "register/register_custom_pass.h"
84+#include "engines/custom_engine/custom_graph_optimizer.h"
85+#include "engines/custom_engine/custom_ops_kernel_info_store.h"
80 86 
81namespace ge {87namespace ge {
82namespace {88namespace {
@@ -117,6 +123,15 @@ struct DummyCompileInfo {
117const std::string kStHostCpuEngine = "DNN_VM_HOST_CPU";123const std::string kStHostCpuEngine = "DNN_VM_HOST_CPU";
118const std::string kStHostCpuKernelStore = "DNN_VM_HOST_CPU_OP_STORE";124const std::string kStHostCpuKernelStore = "DNN_VM_HOST_CPU_OP_STORE";
119 125 
126+class StHostCpuPassCustomOp final : public HostCpuExecuteOp {
127+ public:
128+ graphStatus Execute(gert::HostCpuOpExecutionContext *) override {
129+ return GRAPH_SUCCESS;
130+ }
131+};
132+ 
133+class StDeviceCustomOp final : public BaseCustomOp {};
134+ 
120class FakeUnsupportedHostCpuOpsKernelInfoStore : public OpsKernelInfoStore {135class FakeUnsupportedHostCpuOpsKernelInfoStore : public OpsKernelInfoStore {
121 public:136 public:
122 Status Initialize(const std::map<std::string, std::string> &options) override {137 Status Initialize(const std::map<std::string, std::string> &options) override {
@@ -269,6 +284,24 @@ ComputeGraphPtr BuildHostInputWithoutConsumerGraphForSt() {
269 return graph;284 return graph;
270}285}
271 286 
287+ComputeGraphPtr BuildHostCpuCustomPropagationGraph(const std::string &op_type) {
288+ DEF_GRAPH(g1) {
289+ CHAIN(NODE("data", "Data")->NODE("host_cpu_custom", op_type)->NODE("netoutput", "NetOutput"));
290+ };
291+ 
292+ auto graph = ToComputeGraph(g1);
293+ graph->SetGraphUnknownFlag(true);
294+ graph->FindNode("data")->GetOpDesc()->SetOpKernelLibName(kEngineNameGeLocal);
295+ auto custom_node = graph->FindNode("host_cpu_custom");
296+ custom_node->GetOpDesc()->SetOpEngineName(kEngineNameCustom);
297+ custom_node->GetOpDesc()->SetOpKernelLibName(kCustomOpKernelLibName);
298+ *custom_node->GetOpDesc()->MutableInputDesc(0) = GeTensorDesc(GeShape({1}), FORMAT_ND, DT_INT32);
299+ *custom_node->GetOpDesc()->MutableOutputDesc(0) = GeTensorDesc(GeShape({1}), FORMAT_ND, DT_INT32);
300+ graph->FindNode("netoutput")->GetOpDesc()->SetOpKernelLibName(kEngineNameGeLocal);
301+ (void)AttrUtils::SetBool(graph->FindNode("data")->GetOpDesc(), ATTR_NAME_HOST_TENSOR, true);
302+ return graph;
303+}
304+ 
272template <typename T, typename std::enable_if<(!std::is_array<T>::value), int>::type = 0>305template <typename T, typename std::enable_if<(!std::is_array<T>::value), int>::type = 0>
273static void *CreateCompileInfo() {306static void *CreateCompileInfo() {
274 return new T();307 return new T();
@@ -2051,6 +2084,157 @@ TEST_F(DynamicGraphTest, HostCpuPassDoesNotMarkHostInputWithoutConsumer) {
2051 unsetenv("ENABLE_RUNTIME_V2");2084 unsetenv("ENABLE_RUNTIME_V2");
2052}2085}
2053 2086 
2087+TEST_F(DynamicGraphTest, HostCpuPassMarksHostCpuCustomOp) {
2088+ const AscendString op_type("StHostCpuPassCustomOp");
2089+ CustomOpFactory::RemoveCustomOps({op_type});
2090+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2091+ op_type, OpBackend::kDevice,
2092+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StDeviceCustomOp>(); }),
2093+ GRAPH_SUCCESS);
2094+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2095+ op_type, OpBackend::kHostCPU,
2096+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StHostCpuPassCustomOp>(); }),
2097+ GRAPH_SUCCESS);
2098+ setenv("ENABLE_RUNTIME_V2", "1", 1);
2099+ auto graph = BuildHostCpuCustomPropagationGraph(op_type.GetString());
2100+ ASSERT_NE(graph, nullptr);
2101+ HostcpuEngineUpdatePass pass;
2102+ NodeEngineMap node_atomic_engine_map;
2103+ NodeEngineMap node_composite_engine_map;
2104+ EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS);
2105+ auto custom_node = graph->FindNode("host_cpu_custom");
2106+ ASSERT_NE(custom_node, nullptr);
2107+ EXPECT_EQ(custom_node->GetOpDesc()->GetOpEngineName(), kEngineNameCustom);
2108+ EXPECT_EQ(custom_node->GetOpDesc()->GetOpKernelLibName(), kCustomOpKernelLibName);
2109+ std::string lowering_func;
2110+ EXPECT_TRUE(AttrUtils::GetStr(custom_node->GetOpDesc(), kAttrLowingFunc, lowering_func));
2111+ EXPECT_EQ(lowering_func, kHostCpuCustomOpLowerFunc);
2112+ unsetenv("ENABLE_RUNTIME_V2");
2113+ CustomOpFactory::RemoveCustomOps({op_type});
2114+}
2115+ 
2116+TEST_F(DynamicGraphTest, DnnEngineManagerGetsHostCpuCustomEngineName) {
2117+ const AscendString op_type("StDnnHostCpuCustomOp");
2118+ CustomOpFactory::RemoveCustomOps({op_type});
2119+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2120+ op_type, OpBackend::kHostCPU,
2121+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StHostCpuPassCustomOp>(); }),
2122+ GRAPH_SUCCESS);
2123+ 
2124+ auto op_desc = std::make_shared<OpDesc>("host_cpu_custom", op_type.GetString());
2125+ ASSERT_NE(op_desc, nullptr);
2126+ ASSERT_TRUE(AttrUtils::SetStr(op_desc, kAttrLowingFunc, kHostCpuCustomOpLowerFunc));
2127+ OpInfo matched_op_info;
2128+ EXPECT_EQ(DNNEngineManager::GetInstance().GetHostCpuEngineName({}, op_desc, matched_op_info), kEngineNameCustom);
2129+ EXPECT_EQ(matched_op_info.engine, kEngineNameCustom);
2130+ EXPECT_EQ(matched_op_info.opKernelLib, kCustomOpKernelLibName);
2131+ CustomOpFactory::RemoveCustomOps({op_type});
2132+}
2133+ 
2134+TEST_F(DynamicGraphTest, DnnEngineManagerSelectsHostCpuCustomCandidate) {
2135+ const AscendString op_type("StDnnHostCpuCandidateOp");
2136+ CustomOpFactory::RemoveCustomOps({op_type});
2137+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2138+ op_type, OpBackend::kHostCPU,
2139+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StHostCpuPassCustomOp>(); }),
2140+ GRAPH_SUCCESS);
2141+ 
2142+ auto graph = std::make_shared<ComputeGraph>("dnn_host_cpu_candidate");
2143+ auto op_desc = std::make_shared<OpDesc>("host_cpu_custom", op_type.GetString());
2144+ ASSERT_NE(op_desc, nullptr);
2145+ ASSERT_EQ(op_desc->AddOutputDesc(GeTensorDesc(GeShape({1}), FORMAT_ND, DT_FLOAT)), GRAPH_SUCCESS);
2146+ auto node = graph->AddNode(op_desc);
2147+ ASSERT_NE(node, nullptr);
2148+ 
2149+ auto &ops_kernel_manager = OpsKernelManager::GetInstance();
2150+ ops_kernel_manager.ops_kernel_info_[op_type.GetString()] = {
2151+ OpInfo{"DNN_VM_HOST_CPU", "DNN_VM_HOST_CPU_OP_STORE", 0, false, false, false, "", ""}};
2152+ OpInfo matched_op_info;
2153+ EXPECT_EQ(DNNEngineManager::GetInstance().GetDNNEngineName(node, {}, matched_op_info), kEngineNameCustom);
2154+ EXPECT_EQ(matched_op_info.opKernelLib, kCustomOpKernelLibName);
2155+ ops_kernel_manager.ops_kernel_info_.erase(op_type.GetString());
2156+ CustomOpFactory::RemoveCustomOps({op_type});
2157+}
2158+ 
2159+TEST_F(DynamicGraphTest, EnginePartitionerKeepsHostCpuCustomCluster) {
2160+ const AscendString op_type("StPartitionHostCpuCustomOp");
2161+ CustomOpFactory::RemoveCustomOps({op_type});
2162+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2163+ op_type, OpBackend::kHostCPU,
2164+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StHostCpuPassCustomOp>(); }),
2165+ GRAPH_SUCCESS);
2166+ auto graph = BuildHostCpuCustomPropagationGraph(op_type.GetString());
2167+ ASSERT_NE(graph, nullptr);
2168+ auto custom_node = graph->FindNode("host_cpu_custom");
2169+ ASSERT_NE(custom_node, nullptr);
2170+ ASSERT_TRUE(AttrUtils::SetStr(custom_node->GetOpDesc(), kAttrLowingFunc, kHostCpuCustomOpLowerFunc));
2171+ AttrUtils::SetStr(graph, ATTR_NAME_SESSION_GRAPH_ID, "0");
2172+ 
2173+ EnginePartitioner partitioner;
2174+ EXPECT_EQ(partitioner.Partition(graph, EnginePartitioner::Mode::kSecondPartitioning), SUCCESS);
2175+ CustomOpFactory::RemoveCustomOps({op_type});
2176+}
2177+ 
2178+TEST_F(DynamicGraphTest, MarkUnknownStatusForHostCpuCustomOp) {
2179+ const AscendString op_type("StMarkUnknownHostCpuCustomOp");
2180+ CustomOpFactory::RemoveCustomOps({op_type});
2181+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2182+ op_type, OpBackend::kHostCPU,
2183+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StHostCpuPassCustomOp>(); }),
2184+ GRAPH_SUCCESS);
2185+ auto graph = BuildHostCpuCustomPropagationGraph(op_type.GetString());
2186+ ASSERT_NE(graph, nullptr);
2187+ graph->SetGraphUnknownFlag(false);
2188+ auto custom_node = graph->FindNode("host_cpu_custom");
2189+ ASSERT_NE(custom_node, nullptr);
2190+ ASSERT_TRUE(AttrUtils::SetStr(custom_node->GetOpDesc(), kAttrLowingFunc, kHostCpuCustomOpLowerFunc));
2191+ 
2192+ MarkGraphUnknownStatusPass pass;
2193+ EXPECT_EQ(pass.Run(graph), SUCCESS);
2194+ EXPECT_TRUE(graph->GetGraphUnknownFlag());
2195+ CustomOpFactory::RemoveCustomOps({op_type});
2196+}
2197+ 
2198+TEST_F(DynamicGraphTest, HostCpuCustomOpOptimizerSkipsCompile) {
2199+ const AscendString op_type("StHostCpuOptimizerSkipOp");
2200+ CustomOpFactory::RemoveCustomOps({op_type});
2201+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2202+ op_type, OpBackend::kDevice,
2203+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StDeviceCustomOp>(); }),
2204+ GRAPH_SUCCESS);
2205+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2206+ op_type, OpBackend::kHostCPU,
2207+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StHostCpuPassCustomOp>(); }),
2208+ GRAPH_SUCCESS);
2209+ auto graph = std::make_shared<ComputeGraph>("host_cpu_optimizer_skip");
2210+ auto op_desc = std::make_shared<OpDesc>("host_cpu_node", op_type.GetString());
2211+ ASSERT_NE(op_desc, nullptr);
2212+ ASSERT_EQ(op_desc->AddInputDesc(GeTensorDesc(GeShape({1}), FORMAT_ND, DT_FLOAT)), GRAPH_SUCCESS);
2213+ ASSERT_EQ(op_desc->AddOutputDesc(GeTensorDesc(GeShape({1}), FORMAT_ND, DT_FLOAT)), GRAPH_SUCCESS);
2214+ op_desc->SetOpEngineName(kEngineNameCustom);
2215+ op_desc->SetOpKernelLibName(kCustomOpKernelLibName);
2216+ ASSERT_TRUE(AttrUtils::SetStr(op_desc, kAttrLowingFunc, kHostCpuCustomOpLowerFunc));
2217+ ASSERT_NE(graph->AddNode(op_desc), nullptr);
2218+ CustomGraphOptimizer optimizer;
2219+ EXPECT_EQ(optimizer.OptimizeWholeGraph(*graph), SUCCESS);
2220+ CustomOpFactory::RemoveCustomOps({op_type});
2221+}
2222+ 
2223+TEST_F(DynamicGraphTest, CustomOpsKernelInfoStoreChecksDeviceBackend) {
2224+ const AscendString op_type("StCustomStoreDeviceOp");
2225+ CustomOpFactory::RemoveCustomOps({op_type});
2226+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
2227+ op_type, OpBackend::kDevice,
2228+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StDeviceCustomOp>(); }),
2229+ GRAPH_SUCCESS);
2230+ custom::CustomOpsKernelInfoStore store;
2231+ auto op_desc = std::make_shared<OpDesc>("store_node", op_type.GetString());
2232+ ASSERT_NE(op_desc, nullptr);
2233+ std::string reason;
2234+ EXPECT_TRUE(store.CheckSupported(op_desc, reason));
2235+ CustomOpFactory::RemoveCustomOps({op_type});
2236+}
2237+ 
2054TEST_F(DynamicGraphTest, TestHostCpu) {2238TEST_F(DynamicGraphTest, TestHostCpu) {
2055 MmpaStub::GetInstance().SetImpl(std::make_shared<MockMmpa>());2239 MmpaStub::GetInstance().SetImpl(std::make_shared<MockMmpa>());
2056 MockForGenerateTask("DNN_VM_HOST_CPU_OP_STORE", GenerateTaskForHostCpu);2240 MockForGenerateTask("DNN_VM_HOST_CPU_OP_STORE", GenerateTaskForHostCpu);
@@ -24,6 +24,10 @@
24#include "graph/passes/standard_optimize/constant_folding/dimension_compute_pass.h"24#include "graph/passes/standard_optimize/constant_folding/dimension_compute_pass.h"
25#include "graph/passes/standard_optimize/constant_folding/dimension_adjust_pass.h"25#include "graph/passes/standard_optimize/constant_folding/dimension_adjust_pass.h"
26#include "graph/manager/util/graph_optimize_utility.h"26#include "graph/manager/util/graph_optimize_utility.h"
27+#include "graph/custom_op.h"
28+#include "graph/custom_op_factory.h"
29+#include "exe_graph/runtime/host_cpu_op_execution_context.h"
30+#include "securec.h"
27 31 
28#include "ge_graph_dsl/graph_dsl.h"32#include "ge_graph_dsl/graph_dsl.h"
29#include "ge_graph_dsl/assert/graph_assert.h"33#include "ge_graph_dsl/assert/graph_assert.h"
@@ -34,6 +38,24 @@
34using namespace std;38using namespace std;
35using namespace ge;39using namespace ge;
36 40 
41+namespace {
42+class StHostCpuFoldOp final : public HostCpuExecuteOp {
43+ public:
44+ graphStatus Execute(gert::HostCpuOpExecutionContext *ctx) override {
45+ auto *input = ctx->GetInputTensor(0);
46+ GE_ASSERT_NOTNULL(input);
47+ auto *output = ctx->MallocOutputTensor(0, input->GetShape(), input->GetFormat(), input->GetDataType());
48+ GE_ASSERT_NOTNULL(output);
49+ if ((input->GetSize() > 0U) && (input->GetAddr() != nullptr) && (output->GetAddr() != nullptr)) {
50+ if (memcpy_s(output->GetAddr(), output->GetSize(), input->GetAddr(), input->GetSize()) != EOK) {
51+ return GRAPH_FAILED;
52+ }
53+ }
54+ return GRAPH_SUCCESS;
55+ }
56+};
57+} // namespace
58+ 
37const char *ClipByValue = "ClipByValue";59const char *ClipByValue = "ClipByValue";
38class TestClipByValue : public Kernel {60class TestClipByValue : public Kernel {
39 public:61 public:
@@ -733,3 +755,54 @@ TEST_F(ConstantFoldingTest, TestFolding_Ok_IgnoreFoldingWhen) {
733 EXPECT_NE(node, nullptr);755 EXPECT_NE(node, nullptr);
734 };756 };
735}757}
758+ 
759+/**
760+ * constant constant
761+ * | |
762+ * host_cpu_fold_op --> netoutput
763+ * |
764+ * netoutput
765+ */
766+TEST_F(ConstantFoldingTest, HostCpuCustomOpConstantFoldingCopiesTensor) {
767+ const AscendString op_type("StHostCpuFoldOp");
768+ CustomOpFactory::RemoveCustomOps({op_type});
769+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
770+ op_type, OpBackend::kHostCPU,
771+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<StHostCpuFoldOp>(); }),
772+ GRAPH_SUCCESS);
773+ 
774+ GeTensor weight;
775+ std::vector<uint8_t> data{1U, 2U, 3U};
776+ weight.SetData(data);
777+ GeTensorDesc weight_desc;
778+ weight_desc.SetShape(GeShape({3}));
779+ weight_desc.SetOriginShape(GeShape({3}));
780+ weight_desc.SetFormat(FORMAT_ND);
781+ weight_desc.SetDataType(DT_UINT8);
782+ weight.SetTensorDesc(weight_desc);
783+ 
784+ auto constant = OP_CFG(CONSTANT).TensorDesc(FORMAT_ND, DT_UINT8, {3}).Attr<GeTensor>(ATTR_NAME_WEIGHTS, weight);
785+ auto host_cpu_op = OP_CFG("StHostCpuFoldOp").TensorDesc(FORMAT_ND, DT_UINT8, {3});
786+ auto netouput = OP_CFG(NETOUTPUT).TensorDesc(FORMAT_ND, DT_UINT8, {3});
787+ DEF_GRAPH(g1) {
788+ CHAIN(NODE("constant", constant)->EDGE(0, 0)->NODE("host_cpu_op", host_cpu_op));
789+ CHAIN(NODE("host_cpu_op", host_cpu_op)->EDGE(0, 0)->NODE("netoutput", netouput));
790+ };
791+ auto graph = ToGeGraph(g1);
792+ auto compute_graph = GraphUtilsEx::GetComputeGraph(graph);
793+ auto node_ptr = compute_graph->FindNode("host_cpu_op");
794+ ASSERT_NE(node_ptr, nullptr);
795+ 
796+ auto graph_2 = GraphUtilsEx::CreateGraphFromComputeGraph(compute_graph);
797+ map<AscendString, AscendString> options;
798+ Session session(options);
799+ session.AddGraph(2, graph_2, options);
800+ std::vector<InputTensorInfo> inputs;
801+ auto ret = session.BuildGraph(2, inputs);
802+ EXPECT_EQ(ret, SUCCESS);
803+ CHECK_GRAPH(PreRunAfterBuild) {
804+ const auto folded_node = graph->FindNode("host_cpu_op");
805+ EXPECT_EQ(folded_node, nullptr);
806+ };
807+ CustomOpFactory::RemoveCustomOps({op_type});
808+}
@@ -21,6 +21,7 @@
21#include "engines/custom_engine/custom_ops_kernel_builder.h"21#include "engines/custom_engine/custom_ops_kernel_builder.h"
22#include "engines/custom_engine/custom_ops_kernel_info_store.h"22#include "engines/custom_engine/custom_ops_kernel_info_store.h"
23#include "common/checker.h"23#include "common/checker.h"
24+#include "common/ge_common/ge_types.h"
24#include "exe_graph/runtime/annotated_args_context.h"25#include "exe_graph/runtime/annotated_args_context.h"
25#include "graph/compute_graph.h"26#include "graph/compute_graph.h"
26#include "graph/custom_op_factory.h"27#include "graph/custom_op_factory.h"
@@ -985,7 +986,7 @@ TEST_F(UtestCustomOpsKernelInfoStore, RefreshCapturesHostBackendRegisteredOp) {
985 986 
986 auto op_desc = std::make_shared<OpDesc>(kTestOpType, kTestOpType);987 auto op_desc = std::make_shared<OpDesc>(kTestOpType, kTestOpType);
987 std::string reason;988 std::string reason;
988- EXPECT_TRUE(store.CheckSupported(op_desc, reason));989+ EXPECT_FALSE(store.CheckSupported(op_desc, reason));
989}990}
990 991 
991TEST_F(UtestCustomOpsKernelInfoStore, ThreadSafety) {992TEST_F(UtestCustomOpsKernelInfoStore, ThreadSafety) {
@@ -1047,6 +1048,30 @@ TEST_F(UtestCustomOpsKernelInfoStore, OOptimizeWholeGraphConstructsCompileContex
1047 EXPECT_TRUE(g_compile_context_output_called.load());1048 EXPECT_TRUE(g_compile_context_output_called.load());
1048}1049}
1049 1050 
1051+TEST_F(UtestCustomOpsKernelInfoStore, CustomGraphOptimizerSkipsHostCpuCustomCompile) {
1052+ const std::string kTestOpType = "TestHostCpuCustomOp_SkipDeviceCompile";
1053+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
1054+ AscendString(kTestOpType.c_str()), OpBackend::kDevice,
1055+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<MockCompilableCustomOp>(); }),
1056+ GRAPH_SUCCESS);
1057+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
1058+ AscendString(kTestOpType.c_str()), OpBackend::kHostCPU,
1059+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<MockHostCpuCustomOp>(); }),
1060+ GRAPH_SUCCESS);
1061+ 
1062+ auto graph = std::make_shared<ComputeGraph>("host_cpu_custom_compile_graph");
1063+ auto op_desc = std::make_shared<OpDesc>("host_cpu_custom_node", kTestOpType);
1064+ op_desc->SetOpEngineName(kEngineNameCustom);
1065+ op_desc->SetOpKernelLibName(kCustomOpKernelLibName);
1066+ ASSERT_TRUE(AttrUtils::SetStr(op_desc, kAttrLowingFunc, kHostCpuCustomOpLowerFunc));
1067+ ASSERT_NE(graph->AddNode(op_desc), nullptr);
1068+ 
1069+ MockCompilableCustomOp::ResetCompileCount();
1070+ CustomGraphOptimizer optimizer;
1071+ EXPECT_EQ(optimizer.OptimizeWholeGraph(*graph), SUCCESS);
1072+ EXPECT_EQ(MockCompilableCustomOp::GetCompileCount(), 0);
1073+}
1074+ 
1050TEST_F(UtestCustomOpsKernelInfoStore, GenerateTaskDeclaresAnnotatedArgsAndFillsKernelDef) {1075TEST_F(UtestCustomOpsKernelInfoStore, GenerateTaskDeclaresAnnotatedArgsAndFillsKernelDef) {
1051 const std::string kTestOpType = "TestAnnotatedArgsCustomOp_BuilderTest";1076 const std::string kTestOpType = "TestAnnotatedArgsCustomOp_BuilderTest";
1052 auto creator = []() -> std::unique_ptr<BaseCustomOp> {1077 auto creator = []() -> std::unique_ptr<BaseCustomOp> {
@@ -12,6 +12,7 @@
12#include <gmock/gmock.h>12#include <gmock/gmock.h>
13#include <vector>13#include <vector>
14#include <fstream>14#include <fstream>
15+#include <stdexcept>
15 16 
16#include "macro_utils/dt_public_scope.h"17#include "macro_utils/dt_public_scope.h"
17#include "engines/manager/engine_manager/dnnengine_manager.h"18#include "engines/manager/engine_manager/dnnengine_manager.h"
@@ -20,6 +21,10 @@
20#include "framework/common/ge_inner_error_codes.h"21#include "framework/common/ge_inner_error_codes.h"
21#include "common/opskernel/ops_kernel_info_types.h"22#include "common/opskernel/ops_kernel_info_types.h"
22#include "framework/engine/dnnengine.h"23#include "framework/engine/dnnengine.h"
24+#include "common/ge_common/ge_types.h"
25+#include "graph/ascend_string.h"
26+#include "graph/custom_op.h"
27+#include "graph/custom_op_factory.h"
23#include "graph/op_desc.h"28#include "graph/op_desc.h"
24#include "graph/debug/ge_attr_define.h"29#include "graph/debug/ge_attr_define.h"
25#include "engines/manager/opskernel_manager/dnn_ops_kernel_manager.h"30#include "engines/manager/opskernel_manager/dnn_ops_kernel_manager.h"
@@ -64,6 +69,13 @@ class SubOpsKernelInfoStore2 : public OpsKernelInfoStore {
64 bool check_flag_;69 bool check_flag_;
65};70};
66 71 
72+class ThrowingOpsKernelInfoStore final : public SubOpsKernelInfoStore2 {
73+ public:
74+ bool CheckSupported(const NodePtr &, std::string &, CheckSupportFlag &) const override {
75+ throw std::runtime_error("check support failed");
76+ }
77+};
78+ 
67Status SubOpsKernelInfoStore2::Initialize(const std::map<std::string, std::string> &options) {79Status SubOpsKernelInfoStore2::Initialize(const std::map<std::string, std::string> &options) {
68 return FAILED;80 return FAILED;
69}81}
@@ -190,6 +202,24 @@ TEST_F(UtestDnnengineManager, GetHostCpuEngineName) {
190 EXPECT_EQ(matched_op_info.opKernelLib, "DNN_VM_HOST_CPU_OP_STORE");202 EXPECT_EQ(matched_op_info.opKernelLib, "DNN_VM_HOST_CPU_OP_STORE");
191}203}
192 204 
205+TEST_F(UtestDnnengineManager, GetHostCpuEngineNameForCustomOp) {
206+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator("HostCpuCustom", OpBackend::kHostCPU,
207+ []() -> std::unique_ptr<BaseCustomOp> { return nullptr; }),
208+ GRAPH_SUCCESS);
209+ auto &instance = DNNEngineManager::GetInstance();
210+ OpDescPtr op_desc = std::make_shared<OpDesc>("host_cpu_custom", "HostCpuCustom");
211+ ASSERT_NE(op_desc, nullptr);
212+ ASSERT_TRUE(AttrUtils::SetStr(op_desc, kAttrLowingFunc, kHostCpuCustomOpLowerFunc));
213+ 
214+ OpInfo matched_op_info;
215+ EXPECT_EQ(instance.GetHostCpuEngineName({}, op_desc, matched_op_info), kEngineNameCustom);
216+ EXPECT_EQ(matched_op_info.engine, kEngineNameCustom);
217+ EXPECT_EQ(matched_op_info.opKernelLib, kCustomOpKernelLibName);
218+ EXPECT_EQ(op_desc->GetOpEngineName(), kEngineNameCustom);
219+ EXPECT_EQ(op_desc->GetOpKernelLibName(), kCustomOpKernelLibName);
220+ CustomOpFactory::RemoveCustomOps({AscendString("HostCpuCustom")});
221+}
222+ 
193TEST_F(UtestDnnengineManager, ReadJsonFile) {223TEST_F(UtestDnnengineManager, ReadJsonFile) {
194 auto &instance = DNNEngineManager::GetInstance();224 auto &instance = DNNEngineManager::GetInstance();
195 EXPECT_EQ(instance.ReadJsonFile("", nullptr), FAILED);225 EXPECT_EQ(instance.ReadJsonFile("", nullptr), FAILED);
@@ -327,6 +357,85 @@ TEST_F(UtestDnnengineManager, GetDNNEngineName_not_support_dynamic_shape) {
327 EXPECT_EQ(instance.GetDNNEngineName(node), "");357 EXPECT_EQ(instance.GetDNNEngineName(node), "");
328}358}
329 359 
360+TEST_F(UtestDnnengineManager, GetDNNEngineNameHostCpuCustomWithoutLegacyHostCpuCandidate) {
361+ const std::string op_type = "HostCpuCustomWithoutLegacyHostCpuCandidate";
362+ const AscendString op_type_str(op_type.c_str());
363+ CustomOpFactory::RemoveCustomOps({op_type_str});
364+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(op_type_str, OpBackend::kHostCPU,
365+ []() -> std::unique_ptr<BaseCustomOp> { return nullptr; }),
366+ GRAPH_SUCCESS);
367+ 
368+ auto &instance = DNNEngineManager::GetInstance();
369+ auto &okm = OpsKernelManager::GetInstance();
370+ auto graph = std::make_shared<ComputeGraph>("graph");
371+ auto node = UtAddNode(graph, "host_cpu_custom", op_type, 0, 1);
372+ OpInfo custom_op_info;
373+ custom_op_info.engine = kEngineNameCustom;
374+ custom_op_info.opKernelLib = kCustomOpKernelLibName;
375+ okm.ops_kernel_info_[op_type] = {custom_op_info};
376+ auto custom_store = std::make_shared<SubOpsKernelInfoStore2>();
377+ custom_store->check_flag_ = false;
378+ okm.ops_kernel_store_[kCustomOpKernelLibName] = custom_store;
379+ 
380+ OpInfo matched_op_info;
381+ const std::set<std::string> exclude_engines;
382+ EXPECT_EQ(instance.GetDNNEngineName(node, exclude_engines, matched_op_info), kEngineNameCustom);
383+ EXPECT_EQ(node->GetOpDesc()->GetOpEngineName(), kEngineNameCustom);
384+ EXPECT_EQ(node->GetOpDesc()->GetOpKernelLibName(), kCustomOpKernelLibName);
385+ std::string lowering_func;
386+ EXPECT_TRUE(AttrUtils::GetStr(node->GetOpDesc(), kAttrLowingFunc, lowering_func));
387+ EXPECT_EQ(lowering_func, kHostCpuCustomOpLowerFunc);
388+ 
389+ okm.ops_kernel_info_.erase(op_type);
390+ okm.ops_kernel_store_.erase(kCustomOpKernelLibName);
391+ CustomOpFactory::RemoveCustomOps({op_type_str});
392+}
393+ 
394+TEST_F(UtestDnnengineManager, GetDNNEngineNameHostCpuCustomWithHostCpuCandidate) {
395+ const std::string op_type = "HostCpuCustomWithHostCpuCandidate";
396+ const AscendString op_type_str(op_type.c_str());
397+ CustomOpFactory::RemoveCustomOps({op_type_str});
398+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(op_type_str, OpBackend::kHostCPU,
399+ []() -> std::unique_ptr<BaseCustomOp> { return nullptr; }),
400+ GRAPH_SUCCESS);
401+ 
402+ auto &instance = DNNEngineManager::GetInstance();
403+ auto &okm = OpsKernelManager::GetInstance();
404+ auto graph = std::make_shared<ComputeGraph>("graph");
405+ auto node = UtAddNode(graph, "host_cpu_custom", op_type, 0, 1);
406+ OpInfo host_cpu_info;
407+ host_cpu_info.engine = "DNN_VM_HOST_CPU";
408+ host_cpu_info.opKernelLib = "DNN_VM_HOST_CPU_OP_STORE";
409+ okm.ops_kernel_info_[op_type] = {host_cpu_info};
410+ 
411+ OpInfo matched_op_info;
412+ EXPECT_EQ(instance.GetDNNEngineName(node, {}, matched_op_info), kEngineNameCustom);
413+ EXPECT_EQ(matched_op_info.engine, kEngineNameCustom);
414+ EXPECT_EQ(matched_op_info.opKernelLib, kCustomOpKernelLibName);
415+ EXPECT_EQ(node->GetOpDesc()->GetOpEngineName(), kEngineNameCustom);
416+ EXPECT_EQ(node->GetOpDesc()->GetOpKernelLibName(), kCustomOpKernelLibName);
417+ 
418+ okm.ops_kernel_info_.erase(op_type);
419+ CustomOpFactory::RemoveCustomOps({op_type_str});
420+}
421+ 
422+TEST_F(UtestDnnengineManager, GetDNNEngineNameReturnsEmptyWhenCheckSupportedThrows) {
423+ auto &instance = DNNEngineManager::GetInstance();
424+ auto &okm = OpsKernelManager::GetInstance();
425+ auto graph = std::make_shared<ComputeGraph>("graph");
426+ auto node = UtAddNode(graph, "throwing_op", "ThrowingOp", 0, 1);
427+ OpInfo op_info;
428+ op_info.engine = "AIcoreEngine";
429+ op_info.opKernelLib = "throwing_store";
430+ okm.ops_kernel_info_["ThrowingOp"] = {op_info};
431+ okm.ops_kernel_store_["throwing_store"] = std::make_shared<ThrowingOpsKernelInfoStore>();
432+ 
433+ OpInfo matched_op_info;
434+ EXPECT_EQ(instance.GetDNNEngineName(node, {}, matched_op_info), "");
435+ okm.ops_kernel_info_.erase("ThrowingOp");
436+ okm.ops_kernel_store_.erase("throwing_store");
437+}
438+ 
330TEST_F(UtestDnnengineManager, FinalizeNotInitialized) {439TEST_F(UtestDnnengineManager, FinalizeNotInitialized) {
331 auto &instance = DNNEngineManager::GetInstance();440 auto &instance = DNNEngineManager::GetInstance();
332 instance.init_flag_ = false;441 instance.init_flag_ = false;
@@ -22,6 +22,8 @@
22#include "ge_graph_dsl/graph_dsl.h"22#include "ge_graph_dsl/graph_dsl.h"
23#include "graph/utils/tensor_utils.h"23#include "graph/utils/tensor_utils.h"
24#include "graph/operator_factory.h"24#include "graph/operator_factory.h"
25+#include "graph/custom_op_factory.h"
26+#include "graph/custom_op.h"
25#include "graph/operator_reg.h"27#include "graph/operator_reg.h"
26#include "graph/ge_local_context.h"28#include "graph/ge_local_context.h"
27#include "register/op_impl_registry.h"29#include "register/op_impl_registry.h"
@@ -40,6 +42,13 @@
40 42 
41namespace ge {43namespace ge {
42namespace {44namespace {
45+class HostCpuDynamicPartitionOp final : public HostCpuExecuteOp {
46+ public:
47+ graphStatus Execute(gert::HostCpuOpExecutionContext *) override {
48+ return GRAPH_SUCCESS;
49+ }
50+};
51+ 
43// todo 把注册做成stub的庄能力,不影响其他流程52// todo 把注册做成stub的庄能力,不影响其他流程
44IMPL_OP(AddTilingDepend).TilingInputsDataDependency({1});53IMPL_OP(AddTilingDepend).TilingInputsDataDependency({1});
45IMPL_OP(AddTilingDependPlacementHasAicpu)54IMPL_OP(AddTilingDependPlacementHasAicpu)
@@ -166,6 +175,25 @@ TEST_F(UtestDynamicShapePartition, not_single_op_scene_success) {
166 EXPECT_EQ(partitioner.Partition(), SUCCESS);175 EXPECT_EQ(partitioner.Partition(), SUCCESS);
167}176}
168 177 
178+TEST_F(UtestDynamicShapePartition, custom_host_cpu_node_is_unknown_shape) {
179+ const AscendString op_type("DynamicPartitionHostCpuCustomOp");
180+ CustomOpFactory::RemoveCustomOps({op_type});
181+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
182+ op_type, OpBackend::kHostCPU,
183+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<HostCpuDynamicPartitionOp>(); }),
184+ GRAPH_SUCCESS);
185+ auto graph = std::make_shared<ComputeGraph>("root");
186+ auto node = NodeBuilder("host_custom", op_type.GetString()).AddInputDesc({1}).AddOutputDesc({-1}).Build(graph);
187+ node->GetOpDesc()->SetOpEngineName(kEngineNameCustom);
188+ node->GetOpDesc()->SetOpKernelLibName(kCustomOpKernelLibName);
189+ ASSERT_TRUE(AttrUtils::SetStr(node->GetOpDesc(), kAttrLowingFunc, kHostCpuCustomOpLowerFunc));
190+ DynamicShapePartitioner partitioner(graph);
191+ bool is_unknown = false;
192+ EXPECT_EQ(partitioner.IsUnknownShapeNode(node, is_unknown), SUCCESS);
193+ EXPECT_TRUE(is_unknown);
194+ CustomOpFactory::RemoveCustomOps({op_type});
195+}
196+ 
169TEST_F(UtestDynamicShapePartition, TestSingleOpWithSubGraph) {197TEST_F(UtestDynamicShapePartition, TestSingleOpWithSubGraph) {
170 DEF_GRAPH(partitioned_call) {198 DEF_GRAPH(partitioned_call) {
171 auto cond_data = OP_CFG(DATA).InCnt(1).OutCnt(1).Attr(ATTR_NAME_INDEX, 0).TensorDesc(FORMAT_ND, DT_FLOAT, {16});199 auto cond_data = OP_CFG(DATA).InCnt(1).OutCnt(1).Attr(ATTR_NAME_INDEX, 0).TensorDesc(FORMAT_ND, DT_FLOAT, {16});
@@ -34,10 +34,15 @@
34#include "ge/ge_api.h"34#include "ge/ge_api.h"
35#include "macro_utils/dt_public_unscope.h"35#include "macro_utils/dt_public_unscope.h"
36#include "graph/attribute_group/attr_group_shape_env.h"36#include "graph/attribute_group/attr_group_shape_env.h"
37+#include "graph/custom_op_factory.h"
37 38 
38namespace ge {39namespace ge {
39namespace airut {40namespace airut {
40 41 
42+namespace {
43+class PartitionTestCustomOp : public BaseCustomOp {};
44+} // namespace
45+ 
41class GraphBuilder {46class GraphBuilder {
42 public:47 public:
43 explicit GraphBuilder(const std::string &name) {48 explicit GraphBuilder(const std::string &name) {
@@ -615,6 +620,89 @@ TEST_F(UtestGraphPartition, second_partition_graph_with_user_stream_label) {
615 EXPECT_EQ(ge::GELib::GetInstance()->Finalize(), SUCCESS);620 EXPECT_EQ(ge::GELib::GetInstance()->Finalize(), SUCCESS);
616}621}
617 622 
623+TEST_F(UtestGraphPartition, second_partition_keeps_host_custom_cluster_on_custom_engine) {
624+ const std::string op_type = "HostCustomSecondPartitionOp";
625+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
626+ auto creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<PartitionTestCustomOp>(); };
627+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kHostCPU, creator),
628+ GRAPH_SUCCESS);
629+ 
630+ airut::GraphBuilder graph_builder("host_custom_second_partition");
631+ auto data = graph_builder.AddNode("data", DATA, 1, 1, FORMAT_ND, DT_INT32, {16});
632+ auto host_custom = graph_builder.AddNode("host_custom", op_type, 1, 1, FORMAT_ND, DT_INT32, {16});
633+ host_custom->GetOpDesc()->SetOpEngineName(kEngineNameCustom);
634+ host_custom->GetOpDesc()->SetOpKernelLibName(kCustomOpKernelLibName);
635+ ASSERT_TRUE(AttrUtils::SetStr(host_custom->GetOpDesc(), kAttrLowingFunc, kHostCpuCustomOpLowerFunc));
636+ auto legacy_host_cpu = graph_builder.AddNode("legacy_host_cpu", "LegacyHostCpuOp", 1, 1, FORMAT_ND, DT_INT32, {16});
637+ legacy_host_cpu->GetOpDesc()->SetOpEngineName("DNN_VM_HOST_CPU");
638+ legacy_host_cpu->GetOpDesc()->SetOpKernelLibName(kEngineNameHostCpu);
639+ auto net_output = graph_builder.AddNode("net_output", NETOUTPUT, 1, 1, FORMAT_ND, DT_INT32, {16});
640+ graph_builder.AddDataEdge(data, 0, host_custom, 0);
641+ graph_builder.AddDataEdge(host_custom, 0, legacy_host_cpu, 0);
642+ graph_builder.AddDataEdge(legacy_host_cpu, 0, net_output, 0);
643+ 
644+ auto graph = graph_builder.GetGraph();
645+ ASSERT_TRUE(AttrUtils::SetStr(*graph, ATTR_NAME_SESSION_GRAPH_ID, "0"));
646+ EnginePartitioner partitioner;
647+ ASSERT_EQ(partitioner.Partition(graph, EnginePartitioner::Mode::kSecondPartitioning), SUCCESS);
648+ 
649+ bool found_custom_subgraph = false;
650+ bool found_host_cpu_subgraph = false;
651+ for (const auto &subgraph_info : partitioner.GetSubGraphMap().begin()->second) {
652+ if ((subgraph_info->GetEngineName() == kEngineNameCustom) &&
653+ (subgraph_info->GetSubGraph()->FindNode("host_custom") != nullptr)) {
654+ found_custom_subgraph = true;
655+ EXPECT_EQ(subgraph_info->GetSubGraph()->FindNode("legacy_host_cpu"), nullptr);
656+ }
657+ if ((subgraph_info->GetEngineName() == "DNN_VM_HOST_CPU") &&
658+ (subgraph_info->GetSubGraph()->FindNode("legacy_host_cpu") != nullptr)) {
659+ found_host_cpu_subgraph = true;
660+ EXPECT_EQ(subgraph_info->GetSubGraph()->FindNode("host_custom"), nullptr);
661+ }
662+ }
663+ EXPECT_TRUE(found_custom_subgraph);
664+ EXPECT_TRUE(found_host_cpu_subgraph);
665+ EXPECT_EQ(host_custom->GetOpDesc()->GetOpEngineName(), kEngineNameCustom);
666+ EXPECT_EQ(host_custom->GetOpDesc()->GetOpKernelLibName(), kCustomOpKernelLibName);
667+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
668+}
669+ 
670+TEST_F(UtestGraphPartition, second_partition_keeps_device_custom_cluster_on_ai_core) {
671+ const std::string op_type = "DeviceAndHostCustomSecondPartitionOp";
672+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
673+ auto creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<PartitionTestCustomOp>(); };
674+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kDevice, creator),
675+ GRAPH_SUCCESS);
676+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kHostCPU, creator),
677+ GRAPH_SUCCESS);
678+ 
679+ airut::GraphBuilder graph_builder("device_custom_second_partition");
680+ auto data = graph_builder.AddNode("data", DATA, 1, 1, FORMAT_ND, DT_INT32, {16});
681+ auto device_custom = graph_builder.AddNode("device_custom", op_type, 1, 1, FORMAT_ND, DT_INT32, {16});
682+ device_custom->GetOpDesc()->SetOpEngineName(kEngineNameCustom);
683+ device_custom->GetOpDesc()->SetOpKernelLibName(kCustomOpKernelLibName);
684+ auto net_output = graph_builder.AddNode("net_output", NETOUTPUT, 1, 1, FORMAT_ND, DT_INT32, {16});
685+ graph_builder.AddDataEdge(data, 0, device_custom, 0);
686+ graph_builder.AddDataEdge(device_custom, 0, net_output, 0);
687+ 
688+ auto graph = graph_builder.GetGraph();
689+ ASSERT_TRUE(AttrUtils::SetStr(*graph, ATTR_NAME_SESSION_GRAPH_ID, "0"));
690+ EnginePartitioner partitioner;
691+ ASSERT_EQ(partitioner.Partition(graph, EnginePartitioner::Mode::kSecondPartitioning), SUCCESS);
692+ 
693+ bool found_ai_core_subgraph = false;
694+ for (const auto &subgraph_info : partitioner.GetSubGraphMap().begin()->second) {
695+ if ((subgraph_info->GetEngineName() == kEngineNameAiCore) &&
696+ (subgraph_info->GetSubGraph()->FindNode("device_custom") != nullptr)) {
697+ found_ai_core_subgraph = true;
698+ }
699+ }
700+ EXPECT_TRUE(found_ai_core_subgraph);
701+ EXPECT_EQ(device_custom->GetOpDesc()->GetOpEngineName(), kEngineNameCustom);
702+ EXPECT_EQ(device_custom->GetOpDesc()->GetOpKernelLibName(), kCustomOpKernelLibName);
703+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
704+}
705+ 
618TEST_F(UtestGraphPartition, partition_with_graph_stable_topo_bfs2) {706TEST_F(UtestGraphPartition, partition_with_graph_stable_topo_bfs2) {
619 DEF_GRAPH(graph) {707 DEF_GRAPH(graph) {
620 auto data_0 = OP_CFG(DATA).InCnt(1).OutCnt(1).TensorDesc(FORMAT_ND, DT_INT32, {16});708 auto data_0 = OP_CFG(DATA).InCnt(1).OutCnt(1).TensorDesc(FORMAT_ND, DT_INT32, {16});
@@ -17,6 +17,7 @@
17#include "engines/manager/opskernel_manager/dnn_ops_kernel_manager.h"17#include "engines/manager/opskernel_manager/dnn_ops_kernel_manager.h"
18#include "common/opskernel/ops_kernel_info_types.h"18#include "common/opskernel/ops_kernel_info_types.h"
19#include "common/opskernel/ops_kernel_info_store.h"19#include "common/opskernel/ops_kernel_info_store.h"
20+#include "graph/custom_op_factory.h"
20#include "graph/utils/node_utils.h"21#include "graph/utils/node_utils.h"
21#include "mmpa/mmpa_api.h"22#include "mmpa/mmpa_api.h"
22#include "macro_utils/dt_public_unscope.h"23#include "macro_utils/dt_public_unscope.h"
@@ -31,6 +32,23 @@ using namespace testing;
31namespace ge {32namespace ge {
32namespace {33namespace {
33const std::string kHostCpuKernelStore = "DNN_VM_HOST_CPU_OP_STORE";34const std::string kHostCpuKernelStore = "DNN_VM_HOST_CPU_OP_STORE";
35+const std::string kHostCpuEngine = "DNN_VM_HOST_CPU";
36+ 
37+class HostCpuCustomPassMarkOp : public HostCpuExecuteOp {
38+ public:
39+ graphStatus Execute(gert::HostCpuOpExecutionContext *ctx) override {
40+ (void)ctx;
41+ return GRAPH_SUCCESS;
42+ }
43+};
44+ 
45+class DeviceCustomPassMarkOp : public EagerExecuteOp {
46+ public:
47+ graphStatus Execute(gert::EagerOpExecutionContext *ctx) override {
48+ (void)ctx;
49+ return GRAPH_SUCCESS;
50+ }
51+};
34 52 
35class FakeHostCpuOpsKernelInfoStore : public OpsKernelInfoStore {53class FakeHostCpuOpsKernelInfoStore : public OpsKernelInfoStore {
36 public:54 public:
@@ -350,6 +368,49 @@ ComputeGraphPtr BuildHostInputWithoutConsumerGraph() {
350 return graph;368 return graph;
351}369}
352 370 
371+ComputeGraphPtr BuildHostCpuCustomGraph(const std::string &op_type) {
372+ auto graph = std::make_shared<ge::ComputeGraph>("host_cpu_custom_graph");
373+ auto op_desc = std::make_shared<OpDesc>("host_cpu_custom", op_type);
374+ op_desc->SetOpEngineName(kEngineNameCustom);
375+ op_desc->SetOpKernelLibName(kCustomOpKernelLibName);
376+ op_desc->AddInputDesc(GeTensorDesc(GeShape(std::vector<int64_t>{1}), FORMAT_ND, DT_INT32));
377+ op_desc->AddOutputDesc(GeTensorDesc(GeShape(std::vector<int64_t>{1}), FORMAT_ND, DT_INT32));
378+ (void)graph->AddNode(op_desc);
379+ return graph;
380+}
381+ 
382+ComputeGraphPtr BuildHostCpuCustomPropagationGraph(const std::string &op_type) {
383+ DEF_GRAPH(g1) {
384+ CHAIN(NODE("data", "Data")->NODE("host_cpu_custom", op_type)->NODE("netoutput", "NetOutput"));
385+ };
386+ 
387+ auto graph = ToComputeGraph(g1);
388+ graph->SetGraphUnknownFlag(true);
389+ graph->FindNode("data")->GetOpDesc()->SetOpKernelLibName(kEngineNameGeLocal);
390+ graph->FindNode("host_cpu_custom")->GetOpDesc()->SetOpEngineName(kEngineNameCustom);
391+ graph->FindNode("host_cpu_custom")->GetOpDesc()->SetOpKernelLibName(kCustomOpKernelLibName);
392+ graph->FindNode("netoutput")->GetOpDesc()->SetOpKernelLibName(kEngineNameGeLocal);
393+ SetNoStorage(graph->FindNode("host_cpu_custom")->GetOpDesc(), ge::FORMAT_ND, DT_INT32, {1}, {1});
394+ (void)ge::AttrUtils::SetBool(graph->FindNode("data")->GetOpDesc(), ge::ATTR_NAME_HOST_TENSOR, true);
395+ return graph;
396+}
397+ 
398+ComputeGraphPtr BuildPartitionedCallWithHostCpuCustomSubgraph(const std::string &op_type) {
399+ auto graph = std::make_shared<ge::ComputeGraph>("host_cpu_custom_partitionedcall_graph");
400+ auto partition_desc = std::make_shared<OpDesc>("partitionedcall", PARTITIONEDCALL);
401+ partition_desc->SetOpKernelLibName(kEngineNameGeLocal);
402+ partition_desc->RegisterSubgraphIrName("subgraph", SubgraphType::kStatic);
403+ auto partition_node = graph->AddNode(partition_desc);
404+ 
405+ auto sub_graph = BuildHostCpuCustomGraph(op_type);
406+ partition_desc->AddSubgraphName(sub_graph->GetName());
407+ partition_desc->SetSubgraphInstanceName(0, sub_graph->GetName());
408+ sub_graph->SetParentNode(partition_node);
409+ sub_graph->SetParentGraph(graph);
410+ graph->AddSubgraph(sub_graph);
411+ return graph;
412+}
413+ 
353} // namespace414} // namespace
354class UtestHostcpuEngineUpdatePass : public Test {415class UtestHostcpuEngineUpdatePass : public Test {
355 public:416 public:
@@ -525,6 +586,34 @@ TEST_F(UtestHostcpuEngineUpdatePass, HostCpuRejectsUnsupportedDataTypes) {
525 EXPECT_TRUE(output_atomic_engine_map.count(output_graph->FindNode("gather")) == 0U);586 EXPECT_TRUE(output_atomic_engine_map.count(output_graph->FindNode("gather")) == 0U);
526}587}
527 588 
589+TEST_F(UtestHostcpuEngineUpdatePass, DeviceCustomKernelCanSwitchToLegacyHostCpu) {
590+ setenv("ENABLE_RUNTIME_V2", "1", 1);
591+ OpsKernelManager::GetInstance().ops_kernel_store_[kHostCpuKernelStore] =
592+ std::make_shared<FakeHostCpuOpsKernelInfoStore>(true);
593+ const std::string op_type = "Gather";
594+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
595+ auto creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<DeviceCustomPassMarkOp>(); };
596+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kDevice, creator),
597+ GRAPH_SUCCESS);
598+ auto graph = BuildHostCpuSupportCheckGraph();
599+ ASSERT_NE(graph, nullptr);
600+ auto node = graph->FindNode("gather");
601+ ASSERT_NE(node, nullptr);
602+ node->GetOpDesc()->SetOpEngineName(kEngineNameCustom);
603+ node->GetOpDesc()->SetOpKernelLibName(kCustomOpKernelLibName);
604+ 
605+ HostcpuEngineUpdatePass pass;
606+ NodeEngineMap node_atomic_engine_map;
607+ NodeEngineMap node_composite_engine_map;
608+ EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS);
609+ 
610+ EXPECT_EQ(node->GetOpDesc()->GetOpKernelLibName(), kHostCpuKernelStore);
611+ EXPECT_EQ(node->GetOpDesc()->GetOpEngineName(), kHostCpuEngine);
612+ EXPECT_EQ(node_atomic_engine_map[node], kHostCpuEngine);
613+ EXPECT_EQ(node_composite_engine_map[node], kHostCpuEngine);
614+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
615+}
616+ 
528TEST_F(UtestHostcpuEngineUpdatePass, HostInputWithoutConsumerNotMarkedAsModelInput) {617TEST_F(UtestHostcpuEngineUpdatePass, HostInputWithoutConsumerNotMarkedAsModelInput) {
529 setenv("ENABLE_RUNTIME_V2", "1", 1);618 setenv("ENABLE_RUNTIME_V2", "1", 1);
530 auto graph = BuildHostInputWithoutConsumerGraph();619 auto graph = BuildHostInputWithoutConsumerGraph();
@@ -541,6 +630,210 @@ TEST_F(UtestHostcpuEngineUpdatePass, HostInputWithoutConsumerNotMarkedAsModelInp
541 EXPECT_FALSE(is_host_model_input);630 EXPECT_FALSE(is_host_model_input);
542}631}
543 632 
633+TEST_F(UtestHostcpuEngineUpdatePass, HostCpuCustomOpKeepsCustomEngineWithoutHostPropagation) {
634+ setenv("ENABLE_RUNTIME_V2", "1", 1);
635+ std::map<std::string, std::string> ge_options = {{ge::OO_LEVEL, "O2"}};
636+ const std::unordered_map<std::string, OoInfo> &registered_opt_table =
637+ ge::OptionRegistry::GetInstance().GetRegisteredOptTable();
638+ ge::GetThreadLocalContext().GetOo().Initialize(ge_options, registered_opt_table);
639+ 
640+ const std::string op_type = "HostCpuCustomPassMarkOp";
641+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
642+ auto creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<HostCpuCustomPassMarkOp>(); };
643+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kHostCPU, creator),
644+ GRAPH_SUCCESS);
645+ auto graph = BuildHostCpuCustomGraph(op_type);
646+ ASSERT_NE(graph, nullptr);
647+ auto node = graph->FindNode("host_cpu_custom");
648+ ASSERT_NE(node, nullptr);
649+ 
650+ HostcpuEngineUpdatePass pass;
651+ NodeEngineMap node_atomic_engine_map;
652+ NodeEngineMap node_composite_engine_map;
653+ EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS);
654+ 
655+ auto op_desc = node->GetOpDesc();
656+ ASSERT_NE(op_desc, nullptr);
657+ EXPECT_EQ(op_desc->GetOpEngineName(), kEngineNameCustom);
658+ EXPECT_EQ(op_desc->GetOpKernelLibName(), kCustomOpKernelLibName);
659+ std::string lowering_func;
660+ EXPECT_FALSE(AttrUtils::GetStr(op_desc, kAttrLowingFunc, lowering_func));
661+ EXPECT_EQ(node_atomic_engine_map.count(node), 0U);
662+ EXPECT_EQ(node_composite_engine_map.count(node), 0U);
663+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
664+}
665+ 
666+TEST_F(UtestHostcpuEngineUpdatePass, HostCpuCustomOpDoesNotJoinLegacyHostPropagation) {
667+ const std::string op_type = "HostCpuCustomNoLegacyPropagationOp";
668+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
669+ auto creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<HostCpuCustomPassMarkOp>(); };
670+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kHostCPU, creator),
671+ GRAPH_SUCCESS);
672+ auto graph = BuildHostCpuCustomGraph(op_type);
673+ ASSERT_NE(graph, nullptr);
674+ auto node = graph->FindNode("host_cpu_custom");
675+ ASSERT_NE(node, nullptr);
676+ 
677+ HostcpuEngineUpdatePass pass;
678+ NodeEngineMap node_atomic_engine_map;
679+ NodeEngineMap node_composite_engine_map;
680+ EXPECT_EQ(pass.is_node_execute_on_host_.count(node), 0U);
681+ EXPECT_FALSE(pass.CheckAndMarkHostExec(node, node_atomic_engine_map, node_composite_engine_map));
682+ 
683+ std::string lowering_func;
684+ EXPECT_FALSE(AttrUtils::GetStr(node->GetOpDesc(), kAttrLowingFunc, lowering_func));
685+ EXPECT_EQ(pass.host_exe_ops_.count(node), 0U);
686+ ASSERT_EQ(pass.is_node_execute_on_host_.count(node), 1U);
687+ EXPECT_FALSE(pass.is_node_execute_on_host_[node]);
688+ EXPECT_EQ(node_atomic_engine_map.count(node), 0U);
689+ EXPECT_EQ(node_composite_engine_map.count(node), 0U);
690+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
691+}
692+ 
693+TEST_F(UtestHostcpuEngineUpdatePass, HostOnlyCustomOpOnDeviceMarkedWhenHostPropagationMatched) {
694+ setenv("ENABLE_RUNTIME_V2", "1", 1);
695+ std::map<std::string, std::string> ge_options = {{ge::OO_LEVEL, "O2"}};
696+ const std::unordered_map<std::string, OoInfo> &registered_opt_table =
697+ ge::OptionRegistry::GetInstance().GetRegisteredOptTable();
698+ ge::GetThreadLocalContext().GetOo().Initialize(ge_options, registered_opt_table);
699+ 
700+ const std::string op_type = "HostOnlyCustomHostPropagationOp";
701+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
702+ auto creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<HostCpuCustomPassMarkOp>(); };
703+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kHostCPU, creator),
704+ GRAPH_SUCCESS);
705+ auto graph = BuildHostCpuCustomPropagationGraph(op_type);
706+ ASSERT_NE(graph, nullptr);
707+ auto node = graph->FindNode("host_cpu_custom");
708+ ASSERT_NE(node, nullptr);
709+ node->GetOpDesc()->SetOpEngineName(kEngineNameAiCore);
710+ node->GetOpDesc()->SetOpKernelLibName(kEngineNameAiCore);
711+ 
712+ HostcpuEngineUpdatePass pass;
713+ NodeEngineMap node_atomic_engine_map;
714+ NodeEngineMap node_composite_engine_map;
715+ EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS);
716+ 
717+ auto op_desc = node->GetOpDesc();
718+ ASSERT_NE(op_desc, nullptr);
719+ EXPECT_EQ(op_desc->GetOpEngineName(), kEngineNameCustom);
720+ EXPECT_EQ(op_desc->GetOpKernelLibName(), kCustomOpKernelLibName);
721+ std::string lowering_func;
722+ EXPECT_TRUE(AttrUtils::GetStr(op_desc, kAttrLowingFunc, lowering_func));
723+ EXPECT_EQ(lowering_func, kHostCpuCustomOpLowerFunc);
724+ EXPECT_EQ(pass.host_exe_ops_.count(node), 1U);
725+ EXPECT_EQ(node_atomic_engine_map[node], kEngineNameCustom);
726+ EXPECT_EQ(node_composite_engine_map[node], kEngineNameCustom);
727+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
728+}
729+ 
730+TEST_F(UtestHostcpuEngineUpdatePass, DeviceAndHostCustomOpKeepsDevicePathWhenHostPropagationNotMatched) {
731+ setenv("ENABLE_RUNTIME_V2", "1", 1);
732+ std::map<std::string, std::string> ge_options = {{ge::OO_LEVEL, "O2"}};
733+ const std::unordered_map<std::string, OoInfo> &registered_opt_table =
734+ ge::OptionRegistry::GetInstance().GetRegisteredOptTable();
735+ ge::GetThreadLocalContext().GetOo().Initialize(ge_options, registered_opt_table);
736+ 
737+ const std::string op_type = "DeviceAndHostCustomNoHostPropagationOp";
738+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
739+ auto device_creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<DeviceCustomPassMarkOp>(); };
740+ auto host_creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<HostCpuCustomPassMarkOp>(); };
741+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kDevice, device_creator),
742+ GRAPH_SUCCESS);
743+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kHostCPU, host_creator),
744+ GRAPH_SUCCESS);
745+ auto graph = BuildHostCpuCustomGraph(op_type);
746+ ASSERT_NE(graph, nullptr);
747+ auto node = graph->FindNode("host_cpu_custom");
748+ ASSERT_NE(node, nullptr);
749+ 
750+ HostcpuEngineUpdatePass pass;
751+ NodeEngineMap node_atomic_engine_map;
752+ NodeEngineMap node_composite_engine_map;
753+ EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS);
754+ 
755+ auto op_desc = node->GetOpDesc();
756+ ASSERT_NE(op_desc, nullptr);
757+ EXPECT_EQ(op_desc->GetOpEngineName(), kEngineNameCustom);
758+ EXPECT_EQ(op_desc->GetOpKernelLibName(), kCustomOpKernelLibName);
759+ std::string lowering_func;
760+ EXPECT_FALSE(AttrUtils::GetStr(op_desc, kAttrLowingFunc, lowering_func));
761+ EXPECT_EQ(node_atomic_engine_map.count(node), 0U);
762+ EXPECT_EQ(node_composite_engine_map.count(node), 0U);
763+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
764+}
765+ 
766+TEST_F(UtestHostcpuEngineUpdatePass, DeviceAndHostCustomOpMarkedOnlyWhenHostPropagationMatched) {
767+ setenv("ENABLE_RUNTIME_V2", "1", 1);
768+ std::map<std::string, std::string> ge_options = {{ge::OO_LEVEL, "O2"}};
769+ const std::unordered_map<std::string, OoInfo> &registered_opt_table =
770+ ge::OptionRegistry::GetInstance().GetRegisteredOptTable();
771+ ge::GetThreadLocalContext().GetOo().Initialize(ge_options, registered_opt_table);
772+ 
773+ const std::string op_type = "DeviceAndHostCustomHostPropagationOp";
774+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
775+ auto device_creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<DeviceCustomPassMarkOp>(); };
776+ auto host_creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<HostCpuCustomPassMarkOp>(); };
777+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kDevice, device_creator),
778+ GRAPH_SUCCESS);
779+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kHostCPU, host_creator),
780+ GRAPH_SUCCESS);
781+ auto graph = BuildHostCpuCustomPropagationGraph(op_type);
782+ ASSERT_NE(graph, nullptr);
783+ auto node = graph->FindNode("host_cpu_custom");
784+ ASSERT_NE(node, nullptr);
785+ 
786+ HostcpuEngineUpdatePass pass;
787+ NodeEngineMap node_atomic_engine_map;
788+ NodeEngineMap node_composite_engine_map;
789+ EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS);
790+ 
791+ auto op_desc = node->GetOpDesc();
792+ ASSERT_NE(op_desc, nullptr);
793+ EXPECT_EQ(op_desc->GetOpEngineName(), kEngineNameCustom);
794+ EXPECT_EQ(op_desc->GetOpKernelLibName(), kCustomOpKernelLibName);
795+ std::string lowering_func;
796+ EXPECT_TRUE(AttrUtils::GetStr(op_desc, kAttrLowingFunc, lowering_func));
797+ EXPECT_EQ(lowering_func, kHostCpuCustomOpLowerFunc);
798+ EXPECT_EQ(pass.host_exe_ops_.count(node), 1U);
799+ EXPECT_EQ(node_atomic_engine_map[node], kEngineNameCustom);
800+ EXPECT_EQ(node_composite_engine_map[node], kEngineNameCustom);
801+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
802+}
803+ 
804+TEST_F(UtestHostcpuEngineUpdatePass, HostCpuCustomOpInPartitionedCallSubgraphWithoutHostPropagation) {
805+ setenv("ENABLE_RUNTIME_V2", "1", 1);
806+ std::map<std::string, std::string> ge_options = {{ge::OO_LEVEL, "O2"}};
807+ const std::unordered_map<std::string, OoInfo> &registered_opt_table =
808+ ge::OptionRegistry::GetInstance().GetRegisteredOptTable();
809+ ge::GetThreadLocalContext().GetOo().Initialize(ge_options, registered_opt_table);
810+ 
811+ const std::string op_type = "HostCpuCustomPartitionedCallSubgraphOp";
812+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
813+ auto creator = []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<HostCpuCustomPassMarkOp>(); };
814+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(AscendString(op_type.c_str()), OpBackend::kHostCPU, creator),
815+ GRAPH_SUCCESS);
816+ auto graph = BuildPartitionedCallWithHostCpuCustomSubgraph(op_type);
817+ ASSERT_NE(graph, nullptr);
818+ 
819+ HostcpuEngineUpdatePass pass;
820+ NodeEngineMap node_atomic_engine_map;
821+ NodeEngineMap node_composite_engine_map;
822+ EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS);
823+ 
824+ auto partitioned_call = graph->FindNode("partitionedcall");
825+ ASSERT_NE(partitioned_call, nullptr);
826+ auto sub_graph = NodeUtils::GetSubgraph(*partitioned_call, 0);
827+ ASSERT_NE(sub_graph, nullptr);
828+ auto host_custom_node = sub_graph->FindNode("host_cpu_custom");
829+ ASSERT_NE(host_custom_node, nullptr);
830+ std::string lowering_func;
831+ EXPECT_FALSE(AttrUtils::GetStr(host_custom_node->GetOpDesc(), kAttrLowingFunc, lowering_func));
832+ EXPECT_EQ(node_atomic_engine_map.count(host_custom_node), 0U);
833+ EXPECT_EQ(node_composite_engine_map.count(host_custom_node), 0U);
834+ CustomOpFactory::RemoveCustomOps({AscendString(op_type.c_str())});
835+}
836+ 
544TEST_F(UtestHostcpuEngineUpdatePass, CheckInputForHostExec) {837TEST_F(UtestHostcpuEngineUpdatePass, CheckInputForHostExec) {
545 HostcpuEngineUpdatePass pass;838 HostcpuEngineUpdatePass pass;
546 ge::OpDescPtr op_desc = std::make_shared<OpDesc>("mapindex", "MapIndex");839 ge::OpDescPtr op_desc = std::make_shared<OpDesc>("mapindex", "MapIndex");
@@ -17,6 +17,8 @@
17 17 
18#include "common/framework_types_internal.h"18#include "common/framework_types_internal.h"
19#include "common/plugin/ge_make_unique_util.h"19#include "common/plugin/ge_make_unique_util.h"
20+#include "graph/custom_op.h"
21+#include "graph/custom_op_factory.h"
20#include "graph/ge_local_context.h"22#include "graph/ge_local_context.h"
21#include "graph/passes/base_pass.h"23#include "graph/passes/base_pass.h"
22#include "graph/passes/standard_optimize/constant_folding/dimension_compute_pass.h"24#include "graph/passes/standard_optimize/constant_folding/dimension_compute_pass.h"
@@ -25,6 +27,7 @@
25#include "host_kernels/kernel_factory.h"27#include "host_kernels/kernel_factory.h"
26#include "graph/utils/constant_utils.h"28#include "graph/utils/constant_utils.h"
27#include "api/gelib/gelib.h"29#include "api/gelib/gelib.h"
30+#include "securec.h"
28#include "macro_utils/dt_public_unscope.h"31#include "macro_utils/dt_public_unscope.h"
29 32 
30namespace ge {33namespace ge {
@@ -41,6 +44,46 @@ const char *WrongYes2 = "WrongYes2";
41const char *WrongYes3 = "WrongYes3";44const char *WrongYes3 = "WrongYes3";
42const char *WhereDynamic2Static = "WhereDynamic2Static";45const char *WhereDynamic2Static = "WhereDynamic2Static";
43const char *WhereDynamic = "WhereDynamic";46const char *WhereDynamic = "WhereDynamic";
47+const char *HostCustomFold = "HostCustomFold";
48+ 
49+class TestHostCustomFoldOp : public HostCpuExecuteOp {
50+ public:
51+ graphStatus Execute(gert::HostCpuOpExecutionContext *ctx) override {
52+ ++execute_count_;
53+ if (ctx == nullptr) {
54+ return GRAPH_FAILED;
55+ }
56+ auto *input = ctx->GetInputTensor(0);
57+ if (input == nullptr) {
58+ return GRAPH_FAILED;
59+ }
60+ auto *output = ctx->MallocOutputTensor(0, input->GetShape(), input->GetFormat(), input->GetDataType());
61+ if (output == nullptr) {
62+ return GRAPH_FAILED;
63+ }
64+ const auto *data = reinterpret_cast<const uint8_t *>(input->GetAddr());
65+ auto *dst = reinterpret_cast<uint8_t *>(output->GetAddr());
66+ const auto size = input->GetSize();
67+ if ((size > 0U) && (data != nullptr) && (dst != nullptr)) {
68+ (void)memcpy_s(dst, size, data, size);
69+ }
70+ return GRAPH_SUCCESS;
71+ }
72+ 
73+ static int32_t execute_count_;
74+};
75+ 
76+int32_t TestHostCustomFoldOp::execute_count_ = 0;
77+REG_OP_BACKEND(TestHostCustomFoldOp, "HostCustomFold", ge::OpBackend::kHostCPU);
78+ 
79+class TestHostCustomFoldFailureOp final : public HostCpuExecuteOp {
80+ public:
81+ graphStatus Execute(gert::HostCpuOpExecutionContext *) override {
82+ return GRAPH_FAILED;
83+ }
84+};
85+ 
86+class TestNonHostCustomFoldOp final : public BaseCustomOp {};
44 87 
45class TestAddNKernel : public Kernel {88class TestAddNKernel : public Kernel {
46 public:89 public:
@@ -999,6 +1042,149 @@ TEST_F(UtestGraphPassesConstantFoldingPass, testComputeWithHostCpuKernel) {
999 EXPECT_EQ(ret, UNSUPPORTED);1042 EXPECT_EQ(ret, UNSUPPORTED);
1000}1043}
1001 1044 
1045+TEST_F(UtestGraphPassesConstantFoldingPass, test_compute_with_host_cpu_custom_op) {
1046+ TestHostCustomFoldOp::execute_count_ = 0;
1047+ auto builder = ut::GraphBuilder("test");
1048+ auto input = builder.AddNode("input", CONSTANT, 0, 1, FORMAT_NCHW, DT_UINT8, {3});
1049+ auto output = builder.AddNode("output", HostCustomFold, 1, 1, FORMAT_NCHW, DT_UINT8, {3});
1050+ builder.AddDataEdge(input, 0, output, 0);
1051+ auto graph = builder.GetGraph();
1052+ ASSERT_NE(nullptr, graph);
1053+ 
1054+ auto input_tensor = MakeShared<GeTensor>();
1055+ ASSERT_NE(nullptr, input_tensor);
1056+ input_tensor->MutableTensorDesc().SetShape(GeShape({3}));
1057+ input_tensor->MutableTensorDesc().SetOriginShape(GeShape({3}));
1058+ input_tensor->MutableTensorDesc().SetFormat(FORMAT_NCHW);
1059+ input_tensor->MutableTensorDesc().SetOriginFormat(FORMAT_NCHW);
1060+ input_tensor->MutableTensorDesc().SetDataType(DT_UINT8);
1061+ input_tensor->MutableTensorDesc().SetOriginDataType(DT_UINT8);
1062+ std::vector<uint8_t> data{1, 2, 3};
1063+ EXPECT_EQ(input_tensor->SetData(data), SUCCESS);
1064+ ConstantUtils::SetWeight(input->GetOpDesc(), 0, input_tensor);
1065+ 
1066+ ConstantFoldingPass pass;
1067+ std::vector<ConstGeTensorPtr> inputs;
1068+ std::vector<GeTensorPtr> outputs;
1069+ inputs.emplace_back(input_tensor);
1070+ auto ret = pass.ComputeWithHostCpuCustomOp(output, inputs, outputs);
1071+ EXPECT_EQ(ret, SUCCESS);
1072+ EXPECT_EQ(TestHostCustomFoldOp::execute_count_, 1);
1073+ ASSERT_EQ(outputs.size(), 1U);
1074+ EXPECT_EQ(outputs[0]->GetTensorDesc().GetDataType(), DT_UINT8);
1075+ EXPECT_EQ(outputs[0]->GetTensorDesc().GetPlacement(), kPlacementHost);
1076+ // Host CPU output buffers are allocated with 512-byte alignment.
1077+ EXPECT_EQ(outputs[0]->GetData().GetSize(), 512U);
1078+ EXPECT_NE(outputs[0]->GetData().GetData(), nullptr);
1079+ EXPECT_EQ(outputs[0]->GetData().GetData()[0], 1U);
1080+ EXPECT_EQ(outputs[0]->GetData().GetData()[1], 2U);
1081+ EXPECT_EQ(outputs[0]->GetData().GetData()[2], 3U);
1082+}
1083+ 
1084+TEST_F(UtestGraphPassesConstantFoldingPass, test_compute_with_host_cpu_custom_op_unsupported_backend) {
1085+ auto builder = ut::GraphBuilder("test");
1086+ auto input = builder.AddNode("input", CONSTANT, 0, 1, FORMAT_NCHW, DT_UINT8, {3});
1087+ auto output = builder.AddNode("output", "HostCustomFoldNoBackend", 1, 1, FORMAT_NCHW, DT_UINT8, {3});
1088+ builder.AddDataEdge(input, 0, output, 0);
1089+ auto graph = builder.GetGraph();
1090+ ASSERT_NE(nullptr, graph);
1091+ 
1092+ auto input_tensor = MakeShared<GeTensor>();
1093+ ASSERT_NE(nullptr, input_tensor);
1094+ input_tensor->MutableTensorDesc().SetShape(GeShape({3}));
1095+ input_tensor->MutableTensorDesc().SetOriginShape(GeShape({3}));
1096+ input_tensor->MutableTensorDesc().SetFormat(FORMAT_NCHW);
1097+ input_tensor->MutableTensorDesc().SetOriginFormat(FORMAT_NCHW);
1098+ input_tensor->MutableTensorDesc().SetDataType(DT_UINT8);
1099+ input_tensor->MutableTensorDesc().SetOriginDataType(DT_UINT8);
1100+ std::vector<uint8_t> data{1, 2, 3};
1101+ EXPECT_EQ(input_tensor->SetData(data), SUCCESS);
1102+ 
1103+ ConstantFoldingPass pass;
1104+ std::vector<ConstGeTensorPtr> inputs;
1105+ std::vector<GeTensorPtr> outputs;
1106+ inputs.emplace_back(input_tensor);
1107+ auto ret = pass.ComputeWithHostCpuCustomOp(output, inputs, outputs);
1108+ EXPECT_EQ(ret, UNSUPPORTED);
1109+}
1110+ 
1111+TEST_F(UtestGraphPassesConstantFoldingPass, test_compute_with_host_cpu_custom_op_create_failed) {
1112+ const auto op_type = AscendString("HostCustomFoldNoInstance");
1113+ CustomOpFactory::RemoveCustomOps({op_type});
1114+ const auto reg_ret = CustomOpFactory::RegisterCustomOpCreator(
1115+ op_type, OpBackend::kHostCPU, []() -> std::unique_ptr<BaseCustomOp> { return nullptr; });
1116+ EXPECT_EQ(reg_ret, SUCCESS);
1117+ 
1118+ auto builder = ut::GraphBuilder("test");
1119+ auto input = builder.AddNode("input", CONSTANT, 0, 1, FORMAT_NCHW, DT_UINT8, {3});
1120+ auto output = builder.AddNode("output", op_type.GetString(), 1, 1, FORMAT_NCHW, DT_UINT8, {3});
1121+ builder.AddDataEdge(input, 0, output, 0);
1122+ auto graph = builder.GetGraph();
1123+ ASSERT_NE(nullptr, graph);
1124+ 
1125+ auto input_tensor = MakeShared<GeTensor>();
1126+ ASSERT_NE(nullptr, input_tensor);
1127+ input_tensor->MutableTensorDesc().SetShape(GeShape({3}));
1128+ input_tensor->MutableTensorDesc().SetOriginShape(GeShape({3}));
1129+ input_tensor->MutableTensorDesc().SetFormat(FORMAT_NCHW);
1130+ input_tensor->MutableTensorDesc().SetOriginFormat(FORMAT_NCHW);
1131+ input_tensor->MutableTensorDesc().SetDataType(DT_UINT8);
1132+ input_tensor->MutableTensorDesc().SetOriginDataType(DT_UINT8);
1133+ std::vector<uint8_t> data{1, 2, 3};
1134+ EXPECT_EQ(input_tensor->SetData(data), SUCCESS);
1135+ 
1136+ ConstantFoldingPass pass;
1137+ std::vector<ConstGeTensorPtr> inputs;
1138+ std::vector<GeTensorPtr> outputs;
1139+ inputs.emplace_back(input_tensor);
1140+ auto ret = pass.ComputeWithHostCpuCustomOp(output, inputs, outputs);
1141+ EXPECT_EQ(ret, PARAM_INVALID);
1142+ 
1143+ CustomOpFactory::RemoveCustomOps({op_type});
1144+}
1145+ 
1146+TEST_F(UtestGraphPassesConstantFoldingPass, test_compute_with_host_cpu_custom_op_execute_failed) {
1147+ const AscendString op_type("HostCustomFoldExecuteFailed");
1148+ CustomOpFactory::RemoveCustomOps({op_type});
1149+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
1150+ op_type, OpBackend::kHostCPU,
1151+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<TestHostCustomFoldFailureOp>(); }),
1152+ SUCCESS);
1153+ auto builder = ut::GraphBuilder("test");
1154+ auto output = builder.AddNode("output", op_type.GetString(), 1, 1, FORMAT_NCHW, DT_UINT8, {3});
1155+ auto input_tensor = MakeShared<GeTensor>();
1156+ ASSERT_NE(input_tensor, nullptr);
1157+ input_tensor->MutableTensorDesc().SetShape(GeShape({3}));
1158+ input_tensor->MutableTensorDesc().SetDataType(DT_UINT8);
1159+ ASSERT_EQ(input_tensor->SetData(std::vector<uint8_t>{1, 2, 3}), SUCCESS);
1160+ ConstantFoldingPass pass;
1161+ std::vector<ConstGeTensorPtr> inputs{input_tensor};
1162+ std::vector<GeTensorPtr> outputs;
1163+ EXPECT_EQ(pass.ComputeWithHostCpuCustomOp(output, inputs, outputs), PARAM_INVALID);
1164+ CustomOpFactory::RemoveCustomOps({op_type});
1165+}
1166+ 
1167+TEST_F(UtestGraphPassesConstantFoldingPass, test_compute_with_host_cpu_custom_op_rejects_wrong_capability) {
1168+ const AscendString op_type("HostCustomFoldWrongCapability");
1169+ CustomOpFactory::RemoveCustomOps({op_type});
1170+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
1171+ op_type, OpBackend::kHostCPU,
1172+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<TestNonHostCustomFoldOp>(); }),
1173+ SUCCESS);
1174+ auto builder = ut::GraphBuilder("test");
1175+ auto output = builder.AddNode("output", op_type.GetString(), 1, 1, FORMAT_NCHW, DT_UINT8, {3});
1176+ auto input_tensor = MakeShared<GeTensor>();
1177+ ASSERT_NE(input_tensor, nullptr);
1178+ input_tensor->MutableTensorDesc().SetShape(GeShape({3}));
1179+ input_tensor->MutableTensorDesc().SetDataType(DT_UINT8);
1180+ ASSERT_EQ(input_tensor->SetData(std::vector<uint8_t>{1, 2, 3}), SUCCESS);
1181+ ConstantFoldingPass pass;
1182+ std::vector<ConstGeTensorPtr> inputs{input_tensor};
1183+ std::vector<GeTensorPtr> outputs;
1184+ EXPECT_EQ(pass.ComputeWithHostCpuCustomOp(output, inputs, outputs), PARAM_INVALID);
1185+ CustomOpFactory::RemoveCustomOps({op_type});
1186+}
1187+ 
1002TEST_F(UtestGraphPassesConstantFoldingPass, ConstantFoldingAddNSuccess) {1188TEST_F(UtestGraphPassesConstantFoldingPass, ConstantFoldingAddNSuccess) {
1003 GraphOptimizeUtility graph_optimize_utility;1189 GraphOptimizeUtility graph_optimize_utility;
1004 auto graph = BuildGraph1();1190 auto graph = BuildGraph1();
@@ -21,6 +21,8 @@
21#include "graph/debug/ge_attr_define.h"21#include "graph/debug/ge_attr_define.h"
22#include "graph/passes/shape_optimize/mark_graph_unknown_status_pass.h"22#include "graph/passes/shape_optimize/mark_graph_unknown_status_pass.h"
23#include "graph/utils/op_desc_utils.h"23#include "graph/utils/op_desc_utils.h"
24+#include "graph/custom_op_factory.h"
25+#include "graph/custom_op.h"
24#include "graph/passes/pass_manager.h"26#include "graph/passes/pass_manager.h"
25#include "api/gelib/gelib.h"27#include "api/gelib/gelib.h"
26#include "engines/manager/opskernel_manager/dnn_ops_kernel_manager.h"28#include "engines/manager/opskernel_manager/dnn_ops_kernel_manager.h"
@@ -157,4 +159,31 @@ TEST_F(MarkGraphUnknownStatusPassTest, HostCpuEngineNoTilingSuccess) {
157 EXPECT_EQ(ret, SUCCESS);159 EXPECT_EQ(ret, SUCCESS);
158 EXPECT_TRUE(graph_->GetGraphUnknownFlag());160 EXPECT_TRUE(graph_->GetGraphUnknownFlag());
159}161}
162+ 
163+TEST_F(MarkGraphUnknownStatusPassTest, HostCpuCustomOpSuccess) {
164+ const AscendString op_type("MarkHostCpuCustomOp");
165+ CustomOpFactory::RemoveCustomOps({op_type});
166+ class HostCpuOpForMarkUnknown final : public HostCpuExecuteOp {
167+ public:
168+ graphStatus Execute(gert::HostCpuOpExecutionContext *) override {
169+ return GRAPH_SUCCESS;
170+ }
171+ };
172+ ASSERT_EQ(CustomOpFactory::RegisterCustomOpCreator(
173+ op_type, OpBackend::kHostCPU,
174+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<HostCpuOpForMarkUnknown>(); }),
175+ GRAPH_SUCCESS);
176+ auto node1 = NewNode("Op1", DATA_TYPE, 0, 1);
177+ auto node2 = NewNode("Op2", op_type.GetString(), 1, 1);
178+ auto net_output = NewNode("NetOutput", NETOUTPUT, 3, 3);
179+ node2->GetOpDesc()->SetOpEngineName(kEngineNameCustom);
180+ node2->GetOpDesc()->SetOpKernelLibName(kCustomOpKernelLibName);
181+ ASSERT_TRUE(AttrUtils::SetStr(node2->GetOpDesc(), kAttrLowingFunc, kHostCpuCustomOpLowerFunc));
182+ GraphUtils::AddEdge(node1->GetOutDataAnchor(0), node2->GetInDataAnchor(0));
183+ GraphUtils::AddEdge(node2->GetOutDataAnchor(0), net_output->GetInDataAnchor(1));
184+ 
185+ EXPECT_EQ(mark_graph_unknown_status_pass_.Run(graph_), SUCCESS);
186+ EXPECT_TRUE(graph_->GetGraphUnknownFlag());
187+ CustomOpFactory::RemoveCustomOps({op_type});
188+}
160} // namespace ge189} // namespace ge
@@ -127,6 +127,110 @@ TEST_F(CustomNodeConverterUT, custom_op_convert_test) {
127 "success");127 "success");
128}128}
129 129 
130+TEST_F(CustomNodeConverterUT, host_cpu_custom_op_convert_test) {
131+ auto graph = ShareGraph::BuildCustomOpGraph();
132+ auto root_model = GeModelBuilder(graph).BuildGeRootModel();
133+ auto global_data = GlobalDataFaker(root_model).Build();
134+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kInit);
135+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kMain);
136+ bg::LowerConstDataNode(global_data);
137+ std::vector<bg::ValueHolderPtr> shapes;
138+ std::vector<bg::DevMemValueHolderPtr> addrs;
139+ LowerDataNodes(graph, global_data, shapes, addrs);
140+ 
141+ LowerInput add_input = {shapes, addrs, &global_data};
142+ auto custom_op = graph->FindNode("custom_op");
143+ auto ret = LoweringHostCustomNode(custom_op, add_input);
144+ ASSERT_TRUE(ret.result.IsSuccess());
145+ ASSERT_EQ(ret.out_addrs.size(), 1);
146+ ASSERT_EQ(ret.out_shapes.size(), 1);
147+ ASSERT_EQ(ret.out_addrs[0]->GetPlacement(), kOnHost);
148+ 
149+ auto frame = bg::ValueHolder::PopGraphFrame();
150+ ASSERT_NE(frame, nullptr);
151+ auto exe_graph = frame->GetExecuteGraph().get();
152+ ASSERT_NE(exe_graph, nullptr);
153+ ASSERT_NE(ge::ExecuteGraphUtils::FindFirstNodeMatchType(exe_graph, "ExecuteHostCustomOp"), nullptr);
154+ ASSERT_EQ(ge::ExecuteGraphUtils::FindFirstNodeMatchType(exe_graph, "ExecuteCustomOp"), nullptr);
155+ ASSERT_NE(ge::ExecuteGraphUtils::FindFirstNodeMatchType(init_frame_->GetExecuteGraph().get(), "FindHostCpuCustomOp"),
156+ nullptr);
157+}
158+ 
159+TEST_F(CustomNodeConverterUT, host_cpu_custom_op_convert_with_inference_rule_test) {
160+ const std::string rule = R"({"shape":{"inputs":[["s0"],["s1"],["s2"]],"outputs":[["s0","s1","s2"]]}})";
161+ auto graph = ShareGraph::BuildCustomOpGraph();
162+ auto custom_op = graph->FindNode("custom_op");
163+ ASSERT_NE(custom_op, nullptr);
164+ AttrUtils::SetStr(custom_op->GetOpDesc(), ge::ATTR_NAME_INFER_RULE, rule);
165+ 
166+ auto root_model = GeModelBuilder(graph).BuildGeRootModel();
167+ auto global_data = GlobalDataFaker(root_model).Build();
168+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kInit);
169+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kMain);
170+ bg::LowerConstDataNode(global_data);
171+ std::vector<bg::ValueHolderPtr> shapes;
172+ std::vector<bg::DevMemValueHolderPtr> addrs;
173+ LowerDataNodes(graph, global_data, shapes, addrs);
174+ 
175+ LowerInput add_input = {shapes, addrs, &global_data};
176+ auto ret = LoweringHostCustomNode(custom_op, add_input);
177+ ASSERT_TRUE(ret.result.IsSuccess());
178+ 
179+ auto frame = bg::ValueHolder::PopGraphFrame();
180+ ASSERT_NE(frame, nullptr);
181+ auto exe_graph = frame->GetExecuteGraph().get();
182+ ASSERT_NE(exe_graph, nullptr);
183+ ASSERT_NE(ge::ExecuteGraphUtils::FindFirstNodeMatchType(exe_graph, "ExecuteHostCustomOpWithInferShape"), nullptr);
184+ ASSERT_EQ(ge::ExecuteGraphUtils::FindFirstNodeMatchType(exe_graph, "ExecuteHostCustomOp"), nullptr);
185+}
186+ 
187+TEST_F(CustomNodeConverterUT, host_cpu_custom_op_convert_skips_optional_unconnected_input) {
188+ auto graph = ShareGraph::BuildCustomOpGraph();
189+ auto custom_op = graph->FindNode("custom_op");
190+ ASSERT_NE(custom_op, nullptr);
191+ ASSERT_EQ(GraphUtils::RemoveEdge(graph->FindNode("data2")->GetOutDataAnchor(0), custom_op->GetInDataAnchor(2)),
192+ GRAPH_SUCCESS);
193+ 
194+ auto root_model = GeModelBuilder(graph).BuildGeRootModel();
195+ auto global_data = GlobalDataFaker(root_model).Build();
196+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kInit);
197+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kMain);
198+ bg::LowerConstDataNode(global_data);
199+ std::vector<bg::ValueHolderPtr> shapes;
200+ std::vector<bg::DevMemValueHolderPtr> addrs;
201+ for (const auto &name : {"data0", "data1"}) {
202+ auto data_ret = LoweringDataNode(graph->FindNode(name), {{}, {}, &global_data});
203+ ASSERT_TRUE(data_ret.result.IsSuccess());
204+ shapes.emplace_back(data_ret.out_shapes[0]);
205+ addrs.emplace_back(data_ret.out_addrs[0]);
206+ graph->FindNode(name)->GetOpDesc()->SetExtAttr(
207+ "_lowering_result", gert::PlacedLoweringResult(graph->FindNode(name), std::move(data_ret)));
208+ }
209+ LowerInput add_input = {shapes, addrs, &global_data};
210+ auto ret = LoweringHostCustomNode(custom_op, add_input);
211+ ASSERT_TRUE(ret.result.IsSuccess());
212+ ASSERT_EQ(ret.out_addrs.size(), 1U);
213+}
214+ 
215+TEST_F(CustomNodeConverterUT, host_cpu_custom_op_convert_adds_input_guard_dependency) {
216+ auto graph = ShareGraph::BuildCustomOpGraph();
217+ auto root_model = GeModelBuilder(graph).BuildGeRootModel();
218+ auto global_data = GlobalDataFaker(root_model).Build();
219+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kInit);
220+ global_data.SetExternalAllocator(nullptr, ExecuteGraphType::kMain);
221+ bg::LowerConstDataNode(global_data);
222+ std::vector<bg::ValueHolderPtr> shapes;
223+ std::vector<bg::DevMemValueHolderPtr> addrs;
224+ LowerDataNodes(graph, global_data, shapes, addrs);
225+ auto input_guarder = ValueHolder::CreateVoidGuarder("FreeInput", addrs[0], {});
226+ ASSERT_NE(input_guarder, nullptr);
227+ addrs[0]->SetGuarder(input_guarder);
228+ 
229+ LowerInput add_input = {shapes, addrs, &global_data};
230+ auto ret = LoweringHostCustomNode(graph->FindNode("custom_op"), add_input);
231+ ASSERT_TRUE(ret.result.IsSuccess());
232+}
233+ 
130TEST_F(CustomNodeConverterUT, custom_op_convert_with_inference_rule_test) {234TEST_F(CustomNodeConverterUT, custom_op_convert_with_inference_rule_test) {
131 const std::string rule = R"({"shape":{"inputs":[["s0"],["s1"],["s2"]],"outputs":[["s0","s1","s2"]]}})";235 const std::string rule = R"({"shape":{"inputs":[["s0"],["s1"],["s2"]],"outputs":[["s0","s1","s2"]]}})";
132 auto graph = ShareGraph::BuildCustomOpGraph();236 auto graph = ShareGraph::BuildCustomOpGraph();
@@ -8,6 +8,10 @@
8 * See LICENSE in the root of the software repository for the full text of the License.8 * See LICENSE in the root of the software repository for the full text of the License.
9 */9 */
10 10 
11+#include <cstdint>
12+#include <memory>
13+#include <vector>
14+ 
11#include <gtest/gtest.h>15#include <gtest/gtest.h>
12#include "engine/custom/converter/custom_node_converter.h"16#include "engine/custom/converter/custom_node_converter.h"
13#include "common/share_graph.h"17#include "common/share_graph.h"
@@ -34,10 +38,12 @@
34#include "register/kernel_registry_impl.h"38#include "register/kernel_registry_impl.h"
35#include "exe_graph/runtime/extended_kernel_context.h"39#include "exe_graph/runtime/extended_kernel_context.h"
36#include "exe_graph/runtime/eager_op_execution_context.h"40#include "exe_graph/runtime/eager_op_execution_context.h"
41+#include "exe_graph/runtime/host_cpu_op_execution_context.h"
37#include "exe_graph/runtime/storage_shape.h"42#include "exe_graph/runtime/storage_shape.h"
38#include "exe_graph/runtime/gert_tensor_data.h"43#include "exe_graph/runtime/gert_tensor_data.h"
39#include "framework/runtime/args_handler.h"44#include "framework/runtime/args_handler.h"
40#include "kernel/common_kernel_impl/infer_shape.h"45#include "kernel/common_kernel_impl/infer_shape.h"
46+#include "graph_metadef/depends/faker/allocator_faker.h"
41 47 
42using namespace ge;48using namespace ge;
43using namespace gert::bg;49using namespace gert::bg;
@@ -109,6 +115,80 @@ class TestRegistryOnlyCustomOp : public EagerExecuteOp {
109 }115 }
110};116};
111 117 
118+class TestRegistryOnlyHostCpuCustomOp : public HostCpuExecuteOp {
119+ public:
120+ graphStatus Execute(gert::HostCpuOpExecutionContext *ctx) override {
121+ (void)ctx;
122+ return SUCCESS;
123+ }
124+};
125+ 
126+class TestShapeInferCustomOp : public ShapeInferOp {
127+ public:
128+ graphStatus InferShape(gert::InferShapeContext *ctx) override {
129+ (void)ctx;
130+ return GRAPH_SUCCESS;
131+ }
132+ 
133+ graphStatus InferDataType(gert::InferDataTypeContext *ctx) override {
134+ (void)ctx;
135+ return GRAPH_SUCCESS;
136+ }
137+};
138+ 
139+class TestHostCpuCustomOp : public HostCpuExecuteOp {
140+ public:
141+ graphStatus Execute(gert::HostCpuOpExecutionContext *ctx) override {
142+ ++execute_count_;
143+ auto input_tensor = ctx->GetInputTensor(0);
144+ GE_ASSERT_NOTNULL(input_tensor);
145+ GE_ASSERT_TRUE(input_tensor->GetPlacement() == kOnHost);
146+ auto output_tensor = ctx->MallocOutputTensor(0, StorageShape({16}, {16}),
147+ StorageFormat(FORMAT_ND, FORMAT_ND, ExpandDimsType()), DT_FLOAT);
148+ GE_ASSERT_NOTNULL(output_tensor);
149+ return SUCCESS;
150+ }
151+ 
152+ int32_t execute_count_ = 0;
153+};
154+ 
155+class TestHostCpuCustomOpWithInferShape : public HostCpuExecuteOp {
156+ public:
157+ graphStatus Execute(gert::HostCpuOpExecutionContext *ctx) override {
158+ ++execute_count_;
159+ const auto *output_desc = ctx->GetOutputTensor(0);
160+ GE_ASSERT_NOTNULL(output_desc);
161+ GE_ASSERT_TRUE(output_desc->GetStorageShape().GetDimNum() == 1);
162+ GE_ASSERT_TRUE(output_desc->GetStorageShape().GetDim(0) == 32);
163+ auto output_tensor =
164+ ctx->MallocOutputTensor(0, output_desc->GetShape(), output_desc->GetFormat(), output_desc->GetDataType());
165+ GE_ASSERT_NOTNULL(output_tensor);
166+ return SUCCESS;
167+ }
168+ 
169+ int32_t execute_count_ = 0;
170+};
171+ 
172+class HostCpuCustomOpAllocatorFaker : public AllocatorFaker {
173+ public:
174+ HostCpuCustomOpAllocatorFaker() {
175+ SetPlacement(kOnHost);
176+ }
177+ 
178+ TensorData MallocTensorDataFromL1(size_t size) override {
179+ if (size == 0U) {
180+ return TensorData();
181+ }
182+ std::unique_ptr<uint8_t[]> block(new uint8_t[size]);
183+ auto *addr = block.get();
184+ l1_blocks_.emplace_back(std::move(block));
185+ return TensorData(addr, nullptr, size, GetPlacement());
186+ }
187+ 
188+ private:
189+ std::vector<std::unique_ptr<uint8_t[]>> l1_blocks_;
190+};
191+ 
112REG_OP(CustomOp)192REG_OP(CustomOp)
113 .INPUT(x1, TensorType::BasicType())193 .INPUT(x1, TensorType::BasicType())
114 .INPUT(x2, TensorType::BasicType())194 .INPUT(x2, TensorType::BasicType())
@@ -313,6 +393,78 @@ TEST_F(CustomNodeKernelUT, find_custom_op_uses_global_registry) {
313 ASSERT_EQ(*run_context.GetContext<KernelContext>()->GetOutputPointer<BaseCustomOp *>(0), global_op);393 ASSERT_EQ(*run_context.GetContext<KernelContext>()->GetOutputPointer<BaseCustomOp *>(0), global_op);
314}394}
315 395 
396+TEST_F(CustomNodeKernelUT, find_host_cpu_custom_op_uses_model_registry) {
397+ const std::string node_type = "RegistryOnlyHostCpuCustomOpForRt2";
398+ auto custom_op_registry = std::make_shared<CustomOpRegistry>();
399+ ASSERT_EQ(custom_op_registry->RegisterCreator(
400+ node_type.c_str(), OpBackend::kHostCPU,
401+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<TestRegistryOnlyHostCpuCustomOp>(); }),
402+ GRAPH_SUCCESS);
403+ auto *expected_op = custom_op_registry->CreateOrGetCustomOp(node_type.c_str(), OpBackend::kHostCPU);
404+ ASSERT_NE(expected_op, nullptr);
405+ 
406+ auto run_context = BuildKernelRunContext(2, 1);
407+ run_context.value_holder[0].Set(const_cast<char *>(node_type.c_str()), nullptr);
408+ run_context.value_holder[1].Set(custom_op_registry.get(), nullptr);
409+ 
410+ auto find_func = KernelRegistry::GetInstance().FindKernelFuncs("FindHostCpuCustomOp");
411+ ASSERT_NE(find_func, nullptr);
412+ ASSERT_EQ(find_func->run_func(run_context), GRAPH_SUCCESS);
413+ ASSERT_EQ(*run_context.GetContext<KernelContext>()->GetOutputPointer<BaseCustomOp *>(0), expected_op);
414+}
415+ 
416+TEST_F(CustomNodeKernelUT, find_host_cpu_custom_op_does_not_use_device_backend) {
417+ const std::string node_type = "DeviceOnlyCustomOpForHostCpuRt2";
418+ auto custom_op_registry = std::make_shared<CustomOpRegistry>();
419+ ASSERT_EQ(custom_op_registry->RegisterCreator(
420+ node_type.c_str(), OpBackend::kDevice,
421+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<TestRegistryOnlyCustomOp>(); }),
422+ GRAPH_SUCCESS);
423+ 
424+ auto run_context = BuildKernelRunContext(2, 1);
425+ run_context.value_holder[0].Set(const_cast<char *>(node_type.c_str()), nullptr);
426+ run_context.value_holder[1].Set(custom_op_registry.get(), nullptr);
427+ 
428+ auto find_func = KernelRegistry::GetInstance().FindKernelFuncs("FindHostCpuCustomOp");
429+ ASSERT_NE(find_func, nullptr);
430+ ASSERT_NE(find_func->run_func(run_context), GRAPH_SUCCESS);
431+}
432+ 
433+TEST_F(CustomNodeKernelUT, find_custom_shape_infer_op_uses_model_registry) {
434+ const std::string node_type = "ShapeInferCustomOpForRt2";
435+ auto custom_op_registry = std::make_shared<CustomOpRegistry>();
436+ ASSERT_EQ(custom_op_registry->RegisterCreator(
437+ node_type.c_str(), OpBackend::kDevice,
438+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<TestShapeInferCustomOp>(); }),
439+ GRAPH_SUCCESS);
440+ auto *expected_op = custom_op_registry->CreateOrGetCustomOp(node_type.c_str(), OpBackend::kDevice);
441+ ASSERT_NE(expected_op, nullptr);
442+ 
443+ auto run_context = BuildKernelRunContext(2, 1);
444+ run_context.value_holder[0].Set(const_cast<char *>(node_type.c_str()), nullptr);
445+ run_context.value_holder[1].Set(custom_op_registry.get(), nullptr);
446+ auto find_func = KernelRegistry::GetInstance().FindKernelFuncs("FindCustomShapeInferOp");
447+ ASSERT_NE(find_func, nullptr);
448+ ASSERT_EQ(find_func->run_func(run_context), GRAPH_SUCCESS);
449+ EXPECT_EQ(*run_context.GetContext<KernelContext>()->GetOutputPointer<ShapeInferOp *>(0),
450+ CustomOpCast<ShapeInferOp>(expected_op));
451+}
452+ 
453+TEST_F(CustomNodeKernelUT, find_custom_shape_infer_op_fails_for_missing_capability) {
454+ const std::string node_type = "EagerOnlyShapeInferOpForRt2";
455+ auto custom_op_registry = std::make_shared<CustomOpRegistry>();
456+ ASSERT_EQ(custom_op_registry->RegisterCreator(
457+ node_type.c_str(), OpBackend::kDevice,
458+ []() -> std::unique_ptr<BaseCustomOp> { return std::make_unique<TestRegistryOnlyCustomOp>(); }),
459+ GRAPH_SUCCESS);
460+ auto run_context = BuildKernelRunContext(2, 1);
461+ run_context.value_holder[0].Set(const_cast<char *>(node_type.c_str()), nullptr);
462+ run_context.value_holder[1].Set(custom_op_registry.get(), nullptr);
463+ auto find_func = KernelRegistry::GetInstance().FindKernelFuncs("FindCustomShapeInferOp");
464+ ASSERT_NE(find_func, nullptr);
465+ EXPECT_NE(find_func->run_func(run_context), GRAPH_SUCCESS);
466+}
467+ 
316TEST_F(CustomNodeKernelUT, custom_op_with_inference_rule_execute_test) {468TEST_F(CustomNodeKernelUT, custom_op_with_inference_rule_execute_test) {
317 RegisterInferShapeKernels();469 RegisterInferShapeKernels();
318 const std::string rule = R"({"shape":{"inputs":[["s0"],["s1"],["s2"]],"outputs":[["s0"]]}})";470 const std::string rule = R"({"shape":{"inputs":[["s0"],["s1"],["s2"]],"outputs":[["s0"]]}})";
@@ -405,5 +557,113 @@ TEST_F(CustomNodeKernelUT, create_custom_op_outputs_fails_without_args_handler_o
405 // unique_ptr 确保 args_handler 在早退时被 delete,不会泄漏557 // unique_ptr 确保 args_handler 在早退时被 delete,不会泄漏
406 EXPECT_NE(funcs->outputs_creator(nullptr, context), ge::GRAPH_SUCCESS);558 EXPECT_NE(funcs->outputs_creator(nullptr, context), ge::GRAPH_SUCCESS);
407}559}
560+ 
561+TEST_F(CustomNodeKernelUT, create_host_cpu_custom_op_outputs_success) {
562+ auto run_context = KernelRunContextFaker()
563+ .KernelIONum(0U, 1U)
564+ .NodeIoNum(0U, 1U)
565+ .NodeOutputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
566+ .Build();
567+ auto *context = run_context.GetContext<KernelContext>();
568+ ASSERT_NE(context, nullptr);
569+ 
570+ auto funcs = KernelRegistry::GetInstance().FindKernelFuncs("ExecuteHostCustomOp");
571+ ASSERT_NE(funcs, nullptr);
572+ ASSERT_NE(funcs->outputs_creator, nullptr);
573+ ASSERT_EQ(funcs->outputs_creator(nullptr, context), ge::GRAPH_SUCCESS);
574+ EXPECT_NE(context->GetOutputPointer<Tensor>(0), nullptr);
575+}
576+ 
577+TEST_F(CustomNodeKernelUT, execute_host_cpu_custom_op_calls_host_execute) {
578+ Tensor input_tensor = {
579+ {{16}, {16}}, {ge::FORMAT_ND, ge::FORMAT_ND, {}}, kOnHost, ge::DT_FLOAT, reinterpret_cast<void *>(0x1234)};
580+ Tensor output_tensor;
581+ HostCpuCustomOpAllocatorFaker allocator;
582+ TestHostCpuCustomOp custom_op;
583+ auto run_context = KernelRunContextFaker()
584+ .KernelIONum(3U, 1U)
585+ .NodeIoNum(1U, 1U)
586+ .IrInstanceNum({1U})
587+ .NodeInputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
588+ .NodeOutputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
589+ .Inputs({&input_tensor, &allocator, &custom_op})
590+ .Outputs({&output_tensor})
591+ .Build();
592+ auto *context = run_context.GetContext<KernelContext>();
593+ ASSERT_NE(context, nullptr);
594+ 
595+ auto funcs = KernelRegistry::GetInstance().FindKernelFuncs("ExecuteHostCustomOp");
596+ ASSERT_NE(funcs, nullptr);
597+ ASSERT_NE(funcs->run_func, nullptr);
598+ ASSERT_EQ(funcs->run_func(context), ge::GRAPH_SUCCESS);
599+ EXPECT_EQ(custom_op.execute_count_, 1);
600+ EXPECT_EQ(output_tensor.GetPlacement(), kOnHost);
601+ EXPECT_EQ(output_tensor.GetStorageShape(), Shape({16}));
602+ EXPECT_NE(output_tensor.GetAddr(), nullptr);
603+ 
604+ auto trace = funcs->trace_printer(context);
605+ EXPECT_FALSE(trace.empty());
606+ ASSERT_NE(funcs->profiling_info_filler, nullptr);
607+ ProfilingInfoWrapper profiling_info;
608+ EXPECT_EQ(funcs->profiling_info_filler(context, profiling_info), ge::GRAPH_SUCCESS);
609+}
610+ 
611+TEST_F(CustomNodeKernelUT, execute_host_cpu_custom_op_with_infer_shape_uses_template_shape) {
612+ Tensor input_tensor = {
613+ {{16}, {16}}, {ge::FORMAT_ND, ge::FORMAT_ND, {}}, kOnHost, ge::DT_FLOAT, reinterpret_cast<void *>(0x1234)};
614+ Tensor output_tensor(StorageShape(), StorageFormat(FORMAT_ND, FORMAT_ND, ExpandDimsType()), ge::DT_FLOAT);
615+ Tensor template_tensor = {
616+ {{32}, {32}}, {ge::FORMAT_ND, ge::FORMAT_ND, {}}, kOnHost, ge::DT_FLOAT, reinterpret_cast<void *>(0x5678)};
617+ HostCpuCustomOpAllocatorFaker allocator;
618+ TestHostCpuCustomOpWithInferShape custom_op;
619+ auto run_context = KernelRunContextFaker()
620+ .KernelIONum(4U, 1U)
621+ .NodeIoNum(1U, 1U)
622+ .IrInstanceNum({1U})
623+ .NodeInputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
624+ .NodeOutputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
625+ .Inputs({&input_tensor, &allocator, &custom_op, &template_tensor})
626+ .Outputs({&output_tensor})
627+ .Build();
628+ auto *context = run_context.GetContext<KernelContext>();
629+ ASSERT_NE(context, nullptr);
630+ 
631+ auto funcs = KernelRegistry::GetInstance().FindKernelFuncs("ExecuteHostCustomOpWithInferShape");
632+ ASSERT_NE(funcs, nullptr);
633+ ASSERT_NE(funcs->run_func, nullptr);
634+ ASSERT_EQ(funcs->run_func(context), ge::GRAPH_SUCCESS);
635+ EXPECT_EQ(custom_op.execute_count_, 1);
636+ EXPECT_EQ(output_tensor.GetPlacement(), kOnHost);
637+ EXPECT_EQ(output_tensor.GetStorageShape(), Shape({32}));
638+ EXPECT_NE(output_tensor.GetAddr(), nullptr);
639+ 
640+ auto trace = funcs->trace_printer(context);
641+ EXPECT_FALSE(trace.empty());
642+ ASSERT_NE(funcs->profiling_info_filler, nullptr);
643+ ProfilingInfoWrapper profiling_info;
644+ EXPECT_EQ(funcs->profiling_info_filler(context, profiling_info), ge::GRAPH_SUCCESS);
645+}
646+ 
647+TEST_F(CustomNodeKernelUT, execute_host_cpu_custom_op_rejects_non_host_capability) {
648+ Tensor input_tensor = {
649+ {{16}, {16}}, {ge::FORMAT_ND, ge::FORMAT_ND, {}}, kOnHost, ge::DT_FLOAT, reinterpret_cast<void *>(0x1234)};
650+ HostCpuCustomOpAllocatorFaker allocator;
651+ TestRegistryOnlyCustomOp custom_op;
652+ Tensor output_tensor;
653+ auto run_context = KernelRunContextFaker()
654+ .KernelIONum(3U, 1U)
655+ .NodeIoNum(1U, 1U)
656+ .IrInstanceNum({1U})
657+ .NodeInputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
658+ .NodeOutputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
659+ .Inputs({&input_tensor, &allocator, &custom_op})
660+ .Outputs({&output_tensor})
661+ .Build();
662+ auto *context = run_context.GetContext<KernelContext>();
663+ ASSERT_NE(context, nullptr);
664+ auto funcs = KernelRegistry::GetInstance().FindKernelFuncs("ExecuteHostCustomOp");
665+ ASSERT_NE(funcs, nullptr);
666+ EXPECT_NE(funcs->run_func(context), ge::GRAPH_SUCCESS);
667+}
408} // namespace kernel668} // namespace kernel
409} // namespace gert669} // namespace gert
@@ -371,7 +371,8 @@ TEST_F(BgInferShapeUT, InferCustomOpShapeWithoutRule) {
371 ASSERT_EQ(out_shapes.size(), 1);371 ASSERT_EQ(out_shapes.size(), 1);
372 ASSERT_EQ(out_shapes[0]->GetFastNode()->GetType(), "InferShape");372 ASSERT_EQ(out_shapes[0]->GetFastNode()->GetType(), "InferShape");
373 373 
374- auto find_node = ge::ExecuteGraphUtils::FindFirstNodeMatchType(init_frame_->GetExecuteGraph().get(), "FindCustomOp");374+ auto find_node =
375+ ge::ExecuteGraphUtils::FindFirstNodeMatchType(init_frame_->GetExecuteGraph().get(), "FindCustomShapeInferOp");
375 ASSERT_NE(find_node, nullptr);376 ASSERT_NE(find_node, nullptr);
376 377 
377 auto main_frame = ValueHolder::PopGraphFrame({}, {}, "NetOutput");378 auto main_frame = ValueHolder::PopGraphFrame({}, {}, "NetOutput");
@@ -408,8 +409,9 @@ TEST_F(BgInferShapeUT, InferCustomOpShapeUsesModelRegistry) {
408 custom_op, {data0_ret.out_shapes[0], data1_ret.out_shapes[0], data2_ret.out_shapes[0]}, global_data);409 custom_op, {data0_ret.out_shapes[0], data1_ret.out_shapes[0], data2_ret.out_shapes[0]}, global_data);
409 ASSERT_EQ(out_shapes.size(), 1);410 ASSERT_EQ(out_shapes.size(), 1);
410 ASSERT_EQ(out_shapes[0]->GetFastNode()->GetType(), "InferShape");411 ASSERT_EQ(out_shapes[0]->GetFastNode()->GetType(), "InferShape");
411- ASSERT_NE(ge::ExecuteGraphUtils::FindFirstNodeMatchType(init_frame_->GetExecuteGraph().get(), "FindCustomOp"),412+ ASSERT_NE(
412- nullptr);413+ ge::ExecuteGraphUtils::FindFirstNodeMatchType(init_frame_->GetExecuteGraph().get(), "FindCustomShapeInferOp"),
414+ nullptr);
413 ASSERT_EQ(ge::ExecuteGraphUtils::FindFirstNodeMatchType(init_frame_->GetExecuteGraph().get(), "FindInferShapeFunc"),415 ASSERT_EQ(ge::ExecuteGraphUtils::FindFirstNodeMatchType(init_frame_->GetExecuteGraph().get(), "FindInferShapeFunc"),
414 nullptr);416 nullptr);
415 417 
@@ -49,8 +49,6 @@ class Rt2CustomShapeInferOp : public ge::ShapeInferOp {
49 }49 }
50};50};
51 51 
52-class Rt2CustomNoShapeInferOp : public ge::BaseCustomOp {};
53- 
54ge::graphStatus CopyInferShape(InferShapeContext *context) {52ge::graphStatus CopyInferShape(InferShapeContext *context) {
55 auto input = context->GetInputShape(0);53 auto input = context->GetInputShape(0);
56 auto output = context->GetOutputShape(0);54 auto output = context->GetOutputShape(0);
@@ -251,7 +249,6 @@ TEST_F(InferShapeKernelTest, infer_shape_uses_input_custom_op) {
251}249}
252 250 
253TEST_F(InferShapeKernelTest, infer_shape_fails_when_custom_op_has_no_shape_infer) {251TEST_F(InferShapeKernelTest, infer_shape_fails_when_custom_op_has_no_shape_infer) {
254- Rt2CustomNoShapeInferOp custom_op;
255 StorageShape input{{2, 3, 4}, {2, 3, 4}};252 StorageShape input{{2, 3, 4}, {2, 3, 4}};
256 Tensor output;253 Tensor output;
257 auto infer_shape_func = kernel::InferCustomOpShapeFromInput;254 auto infer_shape_func = kernel::InferCustomOpShapeFromInput;
@@ -260,7 +257,7 @@ TEST_F(InferShapeKernelTest, infer_shape_fails_when_custom_op_has_no_shape_infer
260 .KernelIONum(3, 1)257 .KernelIONum(3, 1)
261 .NodeIoNum(1, 1)258 .NodeIoNum(1, 1)
262 .NodeOutputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)259 .NodeOutputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
263- .Inputs({&input, static_cast<ge::BaseCustomOp *>(&custom_op), reinterpret_cast<void *>(infer_shape_func)})260+ .Inputs({&input, static_cast<ge::ShapeInferOp *>(nullptr), reinterpret_cast<void *>(infer_shape_func)})
264 .Outputs({&output})261 .Outputs({&output})
265 .Build();262 .Build();
266 263 
@@ -417,7 +417,6 @@ SKIP_METHODS = [
417 "GetAllCustomOpApiSoPaths",417 "GetAllCustomOpApiSoPaths",
418 "CallInitFunc",418 "CallInitFunc",
419 "UpdateFormatImpl",419 "UpdateFormatImpl",
420- "GetGlobalRegistry",
421 "CreateOrGetCustomOpLocked",420 "CreateOrGetCustomOpLocked",
422 "CallInferFuncV1",421 "CallInferFuncV1",
423 "CallInferFuncV2",422 "CallInferFuncV2",
@@ -425,6 +424,7 @@ SKIP_METHODS = [
425 "CallInferFormatFuncV1",424 "CallInferFormatFuncV1",
426 "CallInferFormatFuncV2",425 "CallInferFormatFuncV2",
427 "InferCustomOpShape",426 "InferCustomOpShape",
427+ "GetRealInNodesAndIndex",
428]428]
429 429 
430"""430"""