forked from huawei/mindspore2022
!19264 [MS][LITE][STABLE]support custom kernel split to cpu subgraph
Merge pull request !19264 from chenjianping/kernel_reg
This commit is contained in:
commit
bf1143cfaf
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -149,7 +149,7 @@ int KernelRegistry::GetKernel(const std::vector<Tensor *> &in_tensors, const std
|
|||
} else {
|
||||
kernel::KernelDesc desc;
|
||||
KernelKeyToKernelDesc(key, &desc);
|
||||
auto creator = kernel::RegisterKernel::GetCreator(desc, static_cast<const schema::Primitive *>(primitive));
|
||||
auto creator = kernel::RegisterKernel::GetCreator(static_cast<const schema::Primitive *>(primitive), &desc);
|
||||
if (creator == nullptr) {
|
||||
return RET_NOT_SUPPORT;
|
||||
}
|
||||
|
|
@ -160,7 +160,7 @@ int KernelRegistry::GetKernel(const std::vector<Tensor *> &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;
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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<std::mutex> lock(lock_);
|
||||
if (desc.type == schema::PrimitiveType_Custom) {
|
||||
int data_type_index = static_cast<int>(desc.data_type) - kNumberTypeBegin - 1;
|
||||
if (desc->type == schema::PrimitiveType_Custom) {
|
||||
int data_type_index = static_cast<int>(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];
|
||||
|
|
|
|||
|
|
@ -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<std::string, std::unordered_map<std::string, kernel::CreateKernel *>> &kernel_creators() {
|
||||
return kernel_creators_;
|
||||
|
|
|
|||
Loading…
Reference in New Issue