已合并
feat: part2 完善HostCPU自定义算子运行时与常量折叠支持 #4538
duhua创建于 8月24日
feat: part2 完善HostCPU自定义算子运行时与常量折叠支持 #4538
已合并
共 37 个文件变更+2232-120
| @@ -71,6 +71,9 @@ REGISTER_PROF_TYPE(AicpuHostCompute); | |||
| 71 | REGISTER_PROF_TYPE(LaunchMixKernelWithHandle); | 71 | REGISTER_PROF_TYPE(LaunchMixKernelWithHandle); |
| 72 | REGISTER_PROF_TYPE(LaunchMixKernelWithFlag); | 72 | REGISTER_PROF_TYPE(LaunchMixKernelWithFlag); |
| 73 | REGISTER_PROF_TYPE(ExecuteCustomOp); | 73 | REGISTER_PROF_TYPE(ExecuteCustomOp); |
| 74 | +REGISTER_PROF_TYPE(ExecuteCustomOpWithInferShape); | ||
| 75 | +REGISTER_PROF_TYPE(ExecuteHostCustomOp); | ||
| 76 | +REGISTER_PROF_TYPE(ExecuteHostCustomOpWithInferShape); | ||
| 74 | REGISTER_PROF_NON_LAUNCH_TYPE(AICoreUpdateContext); | 77 | REGISTER_PROF_NON_LAUNCH_TYPE(AICoreUpdateContext); |
| 75 | REGISTER_PROF_NON_LAUNCH_TYPE(AICpuUpdateContext); | 78 | REGISTER_PROF_NON_LAUNCH_TYPE(AICpuUpdateContext); |
| 76 | REGISTER_PROF_NON_LAUNCH_TYPE(StaAutoUpdateContext); | 79 | REGISTER_PROF_NON_LAUNCH_TYPE(StaAutoUpdateContext); |
| @@ -476,6 +476,7 @@ target_link_libraries(ge_compiler | |||
| 476 | error_manager | 476 | error_manager |
| 477 | unified_dlog | 477 | unified_dlog |
| 478 | runtime_headers | 478 | runtime_headers |
| 479 | + gert | ||
| 479 | aihac_symbolizer | 480 | aihac_symbolizer |
| 480 | lowering | 481 | lowering |
| 481 | -Wl,--as-needed | 482 | -Wl,--as-needed |
| @@ -17,7 +17,10 @@ | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | + | ||
| 20 | 21 | ||
| 22 | + | ||
| 23 | + | ||
| 21 | 24 | ||
| 22 | 25 | ||
| 23 | 26 | ||
| @@ -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 | + | ||
| 145 | ge::Status AppendCompileTaskIfNeeded(const ge::NodePtr &node, std::vector<CompileTask> &compile_tasks) { | 157 | ge::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> | |||
| 68 | bool CustomOpsKernelInfoStore::CheckSupported(const OpDescPtr &op_desc, std::string &reason) const { | 68 | bool 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 custom | 73 | } // namespace custom |
| 75 | } // namespace ge | 74 | } // namespace ge |
| @@ -18,7 +18,11 @@ | |||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | + | ||
| 21 | 22 | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 22 | 26 | ||
| 23 | 27 | ||
| 24 | 28 | ||
| @@ -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 | } // namespace | 90 | } // namespace |
| 72 | 91 | ||
| 73 | DNNEngineManager::DNNEngineManager() : init_flag_(false) {} | 92 | DNNEngineManager::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 selection | 439 | // 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 | ||
| 611 | std::string DNNEngineManager::GetHostCpuEngineName(const std::vector<OpInfo> &op_infos, const OpDescPtr &op_desc, | 648 | std::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 | 20 | ||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | + | ||
| 24 | + | ||
| 23 | 25 | ||
| 24 | 26 | ||
| 25 | 27 | ||
| @@ -51,6 +53,17 @@ constexpr int64_t kThresholdForMergeAllToUnknownGraph = -1; | |||
| 51 | constexpr int32_t kBase = 10; | 53 | constexpr int32_t kBase = 10; |
| 52 | const std::string kStableRdfsSort = "3"; | 54 | const std::string kStableRdfsSort = "3"; |
| 53 | constexpr char_t const *kOffline = "offline"; | 55 | constexpr 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 | } // namespace | 67 | } // namespace |
| 55 | 68 | ||
| 56 | namespace ge { | 69 | namespace 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 | 20 | ||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | + | ||
| 24 | + | ||
| 23 | 25 | ||
| 24 | 26 | ||
| 25 | 27 | ||
| 26 | 28 | ||
| 29 | + | ||
| 27 | 30 | ||
| 28 | 31 | ||
| 29 | 32 | ||
| @@ -48,11 +51,23 @@ const char_t *const kTaskL2FusionInfo = "_task_L2FusionInfo"; | |||
| 48 | const char_t *const kDataAnchorIndexForLxfusion = "_data_anchor_index_for_lxfusion"; | 51 | const char_t *const kDataAnchorIndexForLxfusion = "_data_anchor_index_for_lxfusion"; |
| 49 | const char_t *const kEnableCvParallel = "_enable_cv_parallel"; | 52 | const char_t *const kEnableCvParallel = "_enable_cv_parallel"; |
| 50 | const char_t *const kVectorEngineName = "VectorEngine"; | 53 | const char_t *const kVectorEngineName = "VectorEngine"; |
| 54 | +const char_t *const kHostCpuEngineName = "DNN_VM_HOST_CPU"; | ||
| 51 | const std::string kStableRdfsSort = "3"; | 55 | const std::string kStableRdfsSort = "3"; |
| 52 | const int32_t kOneGraph = 1; // only one graph | 56 | const int32_t kOneGraph = 1; // only one graph |
| 53 | const int32_t kRankOne = 1; // order of graph list is 0,1,2,3..., 1 means second order | 57 | const int32_t kRankOne = 1; // order of graph list is 0,1,2,3..., 1 means second order |
| 54 | const int32_t kRankZero = 0; // order of graph list is 0,1,2,3..., 0 means first order | 58 | const int32_t kRankZero = 0; // order of graph list is 0,1,2,3..., 0 means first order |
| 55 | const int64_t kOverflowDefaultValue = -1; | 59 | const 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 | + | ||
| 56 | struct DeviceIndex { | 71 | struct 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 | ||
| 79 | std::string GenClusterEngineName(const NodePtr &node, EnginePartitioner::Mode mode, const NodeEngineMap &engine_map) { | 94 | std::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 | 15 | ||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | + | ||
| 18 | 19 | ||
| 19 | 20 | ||
| 20 | 21 | ||
| @@ -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 | + | ||
| 78 | bool IsControlV2Op(const std::string &op_type) { | 87 | bool IsControlV2Op(const std::string &op_type) { |
| 79 | return kControlV2Types.count(op_type) > 0U; | 88 | return kControlV2Types.count(op_type) > 0U; |
| 80 | } | 89 | } |
| 81 | 90 | ||
| 82 | bool IsExecOnDevice(const OpDescPtr &op_desc) { | 91 | bool 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 | ||
| 87 | bool IsConstOp(const OpDescPtr &op_desc) { | 97 | bool 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 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 15 | 18 | ||
| 16 | namespace ge { | 19 | namespace ge { |
| 17 | namespace { | 20 | namespace { |
| 18 | const char *const kOwnerGraphIsUnknown = "OwnerGraphIsUnknown"; | 21 | const char *const kOwnerGraphIsUnknown = "OwnerGraphIsUnknown"; |
| 19 | const char *const kHostCpuEngineName = "DNN_VM_HOST_CPU"; | 22 | const 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 | } // namespace | 33 | } // namespace |
| 21 | 34 | ||
| 22 | Status MarkGraphUnknownStatusPass::Run(ComputeGraphPtr graph) { | 35 | Status 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 | 11 | ||
| 12 | 12 | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 13 | 16 | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 14 | 29 | ||
| 15 | - | ||
| 16 | 30 | ||
| 17 | 31 | ||
| 18 | 32 | ||
| @@ -26,6 +40,155 @@ const int64_t kShapeCalNum = 8; | |||
| 26 | const char *const kKernelLibName = "aicpu_ascend_kernel"; | 40 | const char *const kKernelLibName = "aicpu_ascend_kernel"; |
| 27 | const char *const kOpsFlagClose = "0"; | 41 | const char *const kOpsFlagClose = "0"; |
| 28 | const char *const kPassName = "ConstantFoldingPass"; | 42 | const 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 >d) 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 | } // namespace | 192 | } // namespace |
| 30 | 193 | ||
| 31 | bool ConstantFoldingPass::NeedIgnorePass(const NodePtr &node) { | 194 | bool 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 cpu | 246 | // 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 | + | ||
| 98 | Status ConstantFoldingPass::ComputeWithBuiltInKernel(NodePtr &node, const vector<ConstGeTensorPtr> &inputs, | 304 | Status 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"; | |||
| 105 | const std::string kFFTSGraphLowerFunc = "ffts_graph_lower_func"; | 105 | const std::string kFFTSGraphLowerFunc = "ffts_graph_lower_func"; |
| 106 | const std::string kFFTSStaticGraphLowerFunc = "ffts_static_graph_lower_func"; | 106 | const std::string kFFTSStaticGraphLowerFunc = "ffts_static_graph_lower_func"; |
| 107 | const std::string kFFTSMixL2LowerFunc = "ffts_mix_l2_lower_func"; | 107 | const std::string kFFTSMixL2LowerFunc = "ffts_mix_l2_lower_func"; |
| 108 | +const std::string kHostCpuCustomOpLowerFunc = "host_cpu_custom_op_lower_func"; | ||
| 108 | // runtime2.0 calculate func | 109 | // runtime2.0 calculate func |
| 109 | const std::string kAttrCalcArgsSizeFunc = "_ge_attr_calculate_func"; | 110 | const std::string kAttrCalcArgsSizeFunc = "_ge_attr_calculate_func"; |
| 110 | const std::string kFFTSMixL2CalcFunc = "ffts_mix_l2_calc_func"; | 111 | const 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 输出 index | 82 | * @param output_index 输出 index |
| 83 | * @param input_index 输入 index | 83 | * @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 | + | ||
| 88 | ge::graphStatus BuildInputTensors(const ge::NodePtr &node, const LowerInput &lower_input, | 100 | ge::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 | } // namespace | 166 | } // namespace |
| 124 | 167 | ||
| 125 | LowerResult LoweringCustomNode(const ge::NodePtr &node, const LowerInput &lower_input) { | 168 | LowerResult 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 | } |
| 181 | REGISTER_NODE_CONVERTER_PLACEMENT(ge::kCustomOpKernelLibName.c_str(), kOnDeviceHbm, LoweringCustomNode); | 224 | REGISTER_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 gert | 272 | } // namespace gert |
| @@ -15,5 +15,6 @@ | |||
| 15 | 15 | ||
| 16 | namespace gert { | 16 | namespace gert { |
| 17 | LowerResult LoweringCustomNode(const ge::NodePtr &node, const LowerInput &lower_input); | 17 | LowerResult 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 | 20 | ||
| @@ -20,6 +20,7 @@ | |||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | + | ||
| 23 | 24 | ||
| 24 | 25 | ||
| 25 | 26 | ||
| @@ -29,6 +30,10 @@ namespace kernel { | |||
| 29 | namespace { | 30 | namespace { |
| 30 | // 自定义算子特有的输入,从 AdditionalInputIndex::kNum 开始 | 31 | // 自定义算子特有的输入,从 AdditionalInputIndex::kNum 开始 |
| 31 | enum class CustomOpInput { kFunc = static_cast<uint32_t>(EagerOpExecutionContext::AdditionalInputIndex::kNum), kEnd }; | 32 | enum 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 | ||
| 33 | std::string PrintNodeType(const KernelContext *context) { | 38 | std::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 | + | ||
| 93 | static ge::graphStatus CreateOutputTensors(const ExtendedKernelContext *extended_kernel_context, | 129 | static 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) { | |||
| 188 | ge::graphStatus ExecuteCustomOpWithInferShapeFunc(KernelContext *context) { | 223 | ge::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 | + | ||
| 196 | ge::graphStatus FreeCustomOpWorkspacesFunc(KernelContext *context) { | 270 | ge::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 | + | ||
| 292 | REGISTER_KERNEL(FindCustomOp).RunFunc(FindCustomOpFunc); | 384 | REGISTER_KERNEL(FindCustomOp).RunFunc(FindCustomOpFunc); |
| 293 | REGISTER_KERNEL(ExecuteCustomOp) | 385 | REGISTER_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); |
| 308 | REGISTER_KERNEL(FreeArgsGuarder).RunFunc(FreeArgsGuarderFunc).ConcurrentCriticalSectionKey(kKernelUseMemory); | 400 | REGISTER_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 kernel | 416 | } // namespace kernel |
| 310 | } // namespace gert | 417 | } // namespace gert |
| @@ -19,6 +19,10 @@ namespace kernel { | |||
| 19 | ge::graphStatus FindCustomOpFunc(KernelContext *context); | 19 | ge::graphStatus FindCustomOpFunc(KernelContext *context); |
| 20 | ge::graphStatus ExecuteCustomOpFunc(KernelContext *context); | 20 | ge::graphStatus ExecuteCustomOpFunc(KernelContext *context); |
| 21 | ge::graphStatus FreeCustomOpWorkspacesFunc(KernelContext *context); | 21 | ge::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 kernel | 26 | } // namespace kernel |
| 23 | } // namespace gert | 27 | } // 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 | ||
| 147 | std::vector<ValueHolderPtr> BuildCustomOpInferShapeGraph(const ge::NodePtr &node, | 147 | std::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 | 36 | ||
| 37 | 37 | ||
| 38 | 38 | ||
| 39 | + | ||
| 40 | + | ||
| 39 | 41 | ||
| 40 | 42 | ||
| 41 | 43 | ||
| 42 | 44 | ||
| 45 | + | ||
| 46 | + | ||
| 47 | + | ||
| 48 | + | ||
| 49 | + | ||
| 50 | + | ||
| 43 | 51 | ||
| 44 | using namespace ge; | 52 | using namespace ge; |
| 45 | namespace gert { | 53 | namespace 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 | + | ||
| 46 | class PlaceLoweringResultSystemTest : public bg::BgTest { | 82 | class 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 | + | ||
| 106 | REG_OP(Add) | 156 | REG_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 | } // namespace | 164 | } // 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 | + | ||
| 116 | TEST_F(PlaceLoweringResultSystemTest, H2DRunAfterLaunch_PlacedLoweringResult) { | 296 | TEST_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 | 20 | ||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | - | ||
| 24 | 23 | ||
| 24 | + | ||
| 25 | 25 | ||
| 26 | 26 | ||
| 27 | 27 | ||
| @@ -53,6 +53,7 @@ | |||
| 53 | 53 | ||
| 54 | 54 | ||
| 55 | 55 | ||
| 56 | + | ||
| 56 | 57 | ||
| 57 | 58 | ||
| 58 | 59 | ||
| @@ -69,7 +70,6 @@ | |||
| 69 | 70 | ||
| 70 | 71 | ||
| 71 | 72 | ||
| 72 | - | ||
| 73 | 73 | ||
| 74 | 74 | ||
| 75 | 75 | ||
| @@ -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 ge | 102 | } // namespace ge |
| 96 | 103 | ||
| 97 | namespace ge { | 104 | namespace 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 | - | ||
| 818 | class StRegistryShapeInferOp : public ShapeInferOp { | 785 | class 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 | ||
| 829 | class StRegistryShapeInferOpOther final : public StRegistryShapeInferOp {}; | 796 | class 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 | + | ||
| 831 | class TestBaseCustomOp : public EagerExecuteOp { | 815 | class 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 | ||
| 2194 | TEST_F(CustomOpFactoryStTest, CustomOpRegistryCoversCompatibilityAndBackendPaths) { | 2211 | TEST_F(CustomOpFactoryStTest, CustomOpRegistryCoversCompatibilityAndBackendPaths) { |
| @@ -25,6 +25,8 @@ | |||
| 25 | 25 | ||
| 26 | 26 | ||
| 27 | 27 | ||
| 28 | + | ||
| 29 | + | ||
| 28 | 30 | ||
| 29 | 31 | ||
| 30 | 32 | ||
| @@ -36,6 +38,8 @@ | |||
| 36 | 38 | ||
| 37 | 39 | ||
| 38 | 40 | ||
| 41 | + | ||
| 42 | + | ||
| 39 | 43 | ||
| 40 | 44 | ||
| 41 | 45 | ||
| @@ -77,6 +81,8 @@ | |||
| 77 | 81 | ||
| 78 | 82 | ||
| 79 | 83 | ||
| 84 | + | ||
| 85 | + | ||
| 80 | 86 | ||
| 81 | namespace ge { | 87 | namespace ge { |
| 82 | namespace { | 88 | namespace { |
| @@ -117,6 +123,15 @@ struct DummyCompileInfo { | |||
| 117 | const std::string kStHostCpuEngine = "DNN_VM_HOST_CPU"; | 123 | const std::string kStHostCpuEngine = "DNN_VM_HOST_CPU"; |
| 118 | const std::string kStHostCpuKernelStore = "DNN_VM_HOST_CPU_OP_STORE"; | 124 | const 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 | + | ||
| 120 | class FakeUnsupportedHostCpuOpsKernelInfoStore : public OpsKernelInfoStore { | 135 | class 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 | + | ||
| 272 | template <typename T, typename std::enable_if<(!std::is_array<T>::value), int>::type = 0> | 305 | template <typename T, typename std::enable_if<(!std::is_array<T>::value), int>::type = 0> |
| 273 | static void *CreateCompileInfo() { | 306 | static 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 | + | ||
| 2054 | TEST_F(DynamicGraphTest, TestHostCpu) { | 2238 | TEST_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 | 24 | ||
| 25 | 25 | ||
| 26 | 26 | ||
| 27 | + | ||
| 28 | + | ||
| 29 | + | ||
| 30 | + | ||
| 27 | 31 | ||
| 28 | 32 | ||
| 29 | 33 | ||
| @@ -34,6 +38,24 @@ | |||
| 34 | using namespace std; | 38 | using namespace std; |
| 35 | using namespace ge; | 39 | using 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 | + | ||
| 37 | const char *ClipByValue = "ClipByValue"; | 59 | const char *ClipByValue = "ClipByValue"; |
| 38 | class TestClipByValue : public Kernel { | 60 | class 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 | 21 | ||
| 22 | 22 | ||
| 23 | 23 | ||
| 24 | + | ||
| 24 | 25 | ||
| 25 | 26 | ||
| 26 | 27 | ||
| @@ -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 | ||
| 991 | TEST_F(UtestCustomOpsKernelInfoStore, ThreadSafety) { | 992 | TEST_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 | + | ||
| 1050 | TEST_F(UtestCustomOpsKernelInfoStore, GenerateTaskDeclaresAnnotatedArgsAndFillsKernelDef) { | 1075 | TEST_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 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | + | ||
| 15 | 16 | ||
| 16 | 17 | ||
| 17 | 18 | ||
| @@ -20,6 +21,10 @@ | |||
| 20 | 21 | ||
| 21 | 22 | ||
| 22 | 23 | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 23 | 28 | ||
| 24 | 29 | ||
| 25 | 30 | ||
| @@ -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 | + | ||
| 67 | Status SubOpsKernelInfoStore2::Initialize(const std::map<std::string, std::string> &options) { | 79 | Status 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 | + | ||
| 193 | TEST_F(UtestDnnengineManager, ReadJsonFile) { | 223 | TEST_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 | + | ||
| 330 | TEST_F(UtestDnnengineManager, FinalizeNotInitialized) { | 439 | TEST_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 | 22 | ||
| 23 | 23 | ||
| 24 | 24 | ||
| 25 | + | ||
| 26 | + | ||
| 25 | 27 | ||
| 26 | 28 | ||
| 27 | 29 | ||
| @@ -40,6 +42,13 @@ | |||
| 40 | 42 | ||
| 41 | namespace ge { | 43 | namespace ge { |
| 42 | namespace { | 44 | namespace { |
| 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的庄能力,不影响其他流程 |
| 44 | IMPL_OP(AddTilingDepend).TilingInputsDataDependency({1}); | 53 | IMPL_OP(AddTilingDepend).TilingInputsDataDependency({1}); |
| 45 | IMPL_OP(AddTilingDependPlacementHasAicpu) | 54 | IMPL_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 | + | ||
| 169 | TEST_F(UtestDynamicShapePartition, TestSingleOpWithSubGraph) { | 197 | TEST_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 | 34 | ||
| 35 | 35 | ||
| 36 | 36 | ||
| 37 | + | ||
| 37 | 38 | ||
| 38 | namespace ge { | 39 | namespace ge { |
| 39 | namespace airut { | 40 | namespace airut { |
| 40 | 41 | ||
| 42 | +namespace { | ||
| 43 | +class PartitionTestCustomOp : public BaseCustomOp {}; | ||
| 44 | +} // namespace | ||
| 45 | + | ||
| 41 | class GraphBuilder { | 46 | class 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 | + | ||
| 618 | TEST_F(UtestGraphPartition, partition_with_graph_stable_topo_bfs2) { | 706 | TEST_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 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | + | ||
| 20 | 21 | ||
| 21 | 22 | ||
| 22 | 23 | ||
| @@ -31,6 +32,23 @@ using namespace testing; | |||
| 31 | namespace ge { | 32 | namespace ge { |
| 32 | namespace { | 33 | namespace { |
| 33 | const std::string kHostCpuKernelStore = "DNN_VM_HOST_CPU_OP_STORE"; | 34 | const 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 | ||
| 35 | class FakeHostCpuOpsKernelInfoStore : public OpsKernelInfoStore { | 53 | class 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 | } // namespace | 414 | } // namespace |
| 354 | class UtestHostcpuEngineUpdatePass : public Test { | 415 | class 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 | + | ||
| 528 | TEST_F(UtestHostcpuEngineUpdatePass, HostInputWithoutConsumerNotMarkedAsModelInput) { | 617 | TEST_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> ®istered_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> ®istered_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> ®istered_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> ®istered_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> ®istered_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 | + | ||
| 544 | TEST_F(UtestHostcpuEngineUpdatePass, CheckInputForHostExec) { | 837 | TEST_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 | 18 | ||
| 19 | 19 | ||
| 20 | + | ||
| 21 | + | ||
| 20 | 22 | ||
| 21 | 23 | ||
| 22 | 24 | ||
| @@ -25,6 +27,7 @@ | |||
| 25 | 27 | ||
| 26 | 28 | ||
| 27 | 29 | ||
| 30 | + | ||
| 28 | 31 | ||
| 29 | 32 | ||
| 30 | namespace ge { | 33 | namespace ge { |
| @@ -41,6 +44,46 @@ const char *WrongYes2 = "WrongYes2"; | |||
| 41 | const char *WrongYes3 = "WrongYes3"; | 44 | const char *WrongYes3 = "WrongYes3"; |
| 42 | const char *WhereDynamic2Static = "WhereDynamic2Static"; | 45 | const char *WhereDynamic2Static = "WhereDynamic2Static"; |
| 43 | const char *WhereDynamic = "WhereDynamic"; | 46 | const 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 | ||
| 45 | class TestAddNKernel : public Kernel { | 88 | class 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 | + | ||
| 1002 | TEST_F(UtestGraphPassesConstantFoldingPass, ConstantFoldingAddNSuccess) { | 1188 | TEST_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 | 21 | ||
| 22 | 22 | ||
| 23 | 23 | ||
| 24 | + | ||
| 25 | + | ||
| 24 | 26 | ||
| 25 | 27 | ||
| 26 | 28 | ||
| @@ -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 ge | 189 | } // 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 | + | ||
| 130 | TEST_F(CustomNodeConverterUT, custom_op_convert_with_inference_rule_test) { | 234 | TEST_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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 11 | 15 | ||
| 12 | 16 | ||
| 13 | 17 | ||
| @@ -34,10 +38,12 @@ | |||
| 34 | 38 | ||
| 35 | 39 | ||
| 36 | 40 | ||
| 41 | + | ||
| 37 | 42 | ||
| 38 | 43 | ||
| 39 | 44 | ||
| 40 | 45 | ||
| 46 | + | ||
| 41 | 47 | ||
| 42 | using namespace ge; | 48 | using namespace ge; |
| 43 | using namespace gert::bg; | 49 | using 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 | + | ||
| 112 | REG_OP(CustomOp) | 192 | REG_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 | + | ||
| 316 | TEST_F(CustomNodeKernelUT, custom_op_with_inference_rule_execute_test) { | 468 | TEST_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 kernel | 668 | } // namespace kernel |
| 409 | } // namespace gert | 669 | } // 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 | - | ||
| 54 | ge::graphStatus CopyInferShape(InferShapeContext *context) { | 52 | ge::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 | ||
| 253 | TEST_F(InferShapeKernelTest, infer_shape_fails_when_custom_op_has_no_shape_infer) { | 251 | TEST_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 | """ |