From a16ae6da43f797193e480ae8dc6d86b4a239c1a2 Mon Sep 17 00:00:00 2001 From: chenjianping Date: Fri, 2 Jul 2021 10:35:15 +0800 Subject: [PATCH] support custom op split to cpu subgraph --- mindspore/lite/include/registry/register_kernel.h | 2 +- mindspore/lite/src/kernel_registry.cc | 4 ++-- mindspore/lite/src/lite_kernel.h | 2 +- mindspore/lite/src/registry/register_kernel.cc | 4 ++-- .../lite/src/registry/register_kernel_impl.cc | 15 ++++++++------- .../lite/src/registry/register_kernel_impl.h | 2 +- 6 files changed, 15 insertions(+), 14 deletions(-) diff --git a/mindspore/lite/include/registry/register_kernel.h b/mindspore/lite/include/registry/register_kernel.h index 9194363b22..82a0e1d6bd 100644 --- a/mindspore/lite/include/registry/register_kernel.h +++ b/mindspore/lite/include/registry/register_kernel.h @@ -93,7 +93,7 @@ class MS_API RegisterKernel { /// \param[in] primitive Define the attributes of op. /// /// \return Function pointer to create a kernel. - static CreateKernel GetCreator(const kernel::KernelDesc &desc, const schema::Primitive *primitive); + static CreateKernel GetCreator(const schema::Primitive *primitive, kernel::KernelDesc *desc); }; /// \brief KernelReg Defined registration class of kernel. diff --git a/mindspore/lite/src/kernel_registry.cc b/mindspore/lite/src/kernel_registry.cc index 4625e02cd3..d21f0158f7 100644 --- a/mindspore/lite/src/kernel_registry.cc +++ b/mindspore/lite/src/kernel_registry.cc @@ -149,7 +149,7 @@ int KernelRegistry::GetKernel(const std::vector &in_tensors, const std } else { kernel::KernelDesc desc; KernelKeyToKernelDesc(key, &desc); - auto creator = kernel::RegisterKernel::GetCreator(desc, static_cast(primitive)); + auto creator = kernel::RegisterKernel::GetCreator(static_cast(primitive), &desc); if (creator == nullptr) { return RET_NOT_SUPPORT; } @@ -160,7 +160,7 @@ int KernelRegistry::GetKernel(const std::vector &in_tensors, const std auto *lite_kernel = new (std::nothrow) kernel::LiteKernel(base_kernel); if (lite_kernel != nullptr) { kernel::KernelKey tmp_key = key; - if (tmp_key.provider == kArchCPU) { + if (desc.arch == kArchCPU) { tmp_key.arch = kernel::kCPU; } else { tmp_key.arch = kernel::kCustom; diff --git a/mindspore/lite/src/lite_kernel.h b/mindspore/lite/src/lite_kernel.h index 4097236fb5..44adeeaa47 100644 --- a/mindspore/lite/src/lite_kernel.h +++ b/mindspore/lite/src/lite_kernel.h @@ -37,7 +37,7 @@ #include "include/delegate.h" namespace mindspore::kernel { -enum KERNEL_ARCH { kCPU, kGPU, kAPU, kNPU, kCustom, kDelegate, kKernelArch_MIN = kCPU, kKernelArch_MAX = kDelegate }; +enum KERNEL_ARCH { kCPU, kGPU, kAPU, kNPU, kCustom, kDelegate, kKernelArch_MIN = kCPU, kKernelArch_MAX = kAPU }; static const char *const kBuiltin = "Builtin"; struct KernelKey { diff --git a/mindspore/lite/src/registry/register_kernel.cc b/mindspore/lite/src/registry/register_kernel.cc index 8d2ea5657b..2bdf48c924 100644 --- a/mindspore/lite/src/registry/register_kernel.cc +++ b/mindspore/lite/src/registry/register_kernel.cc @@ -30,8 +30,8 @@ int RegisterKernel::RegKernel(const std::string &arch, const std::string &provid return lite::RegistryKernelImpl::GetInstance()->RegKernel(arch, provider, data_type, op_type, creator); } -CreateKernel RegisterKernel::GetCreator(const kernel::KernelDesc &desc, const schema::Primitive *primitive) { - return lite::RegistryKernelImpl::GetInstance()->GetProviderCreator(desc, primitive); +CreateKernel RegisterKernel::GetCreator(const schema::Primitive *primitive, kernel::KernelDesc *desc) { + return lite::RegistryKernelImpl::GetInstance()->GetProviderCreator(primitive, desc); } } // namespace kernel } // namespace mindspore diff --git a/mindspore/lite/src/registry/register_kernel_impl.cc b/mindspore/lite/src/registry/register_kernel_impl.cc index 0f53ca7053..462efac5b6 100644 --- a/mindspore/lite/src/registry/register_kernel_impl.cc +++ b/mindspore/lite/src/registry/register_kernel_impl.cc @@ -99,12 +99,12 @@ int RegistryKernelImpl::RegKernel(const std::string &arch, const std::string &pr return RET_OK; } -kernel::CreateKernel RegistryKernelImpl::GetProviderCreator(const kernel::KernelDesc &desc, - const schema::Primitive *primitive) { +kernel::CreateKernel RegistryKernelImpl::GetProviderCreator(const schema::Primitive *primitive, + kernel::KernelDesc *desc) { kernel::CreateKernel creator = nullptr; std::unique_lock lock(lock_); - if (desc.type == schema::PrimitiveType_Custom) { - int data_type_index = static_cast(desc.data_type) - kNumberTypeBegin - 1; + if (desc->type == schema::PrimitiveType_Custom) { + int data_type_index = static_cast(desc->data_type) - kNumberTypeBegin - 1; if (data_type_index < 0) { return nullptr; } @@ -117,22 +117,23 @@ kernel::CreateKernel RegistryKernelImpl::GetProviderCreator(const kernel::Kernel return item.second[custom_type] != nullptr && item.second[custom_type][data_type_index] != nullptr; }); if (archs_iter != archs.end()) { + desc->arch = archs_iter->first; return archs_iter->second[custom_type][data_type_index]; } } return nullptr; } - auto index = GetFuncIndex(desc); + auto index = GetFuncIndex(*desc); if (index >= kKernelMaxNum || index < 0) { return nullptr; } for (auto &&item : kernel_creators_) { - if (item.first != desc.provider) { + if (item.first != desc->provider) { continue; } for (auto &&arch_item : item.second) { - if (arch_item.first != desc.arch) { + if (arch_item.first != desc->arch) { continue; } creator = arch_item.second[index]; diff --git a/mindspore/lite/src/registry/register_kernel_impl.h b/mindspore/lite/src/registry/register_kernel_impl.h index 9025c317ae..0f00aec939 100644 --- a/mindspore/lite/src/registry/register_kernel_impl.h +++ b/mindspore/lite/src/registry/register_kernel_impl.h @@ -47,7 +47,7 @@ class RegistryKernelImpl { int RegKernel(const std::string &arch, const std::string &provider, TypeId data_type, int type, kernel::CreateKernel creator); - virtual kernel::CreateKernel GetProviderCreator(const kernel::KernelDesc &desc, const schema::Primitive *primitive); + virtual kernel::CreateKernel GetProviderCreator(const schema::Primitive *primitive, kernel::KernelDesc *desc); const std::map> &kernel_creators() { return kernel_creators_;