| 25 | } |
| 26 | |
| 27 | std::pair<std::vector<const KernelFunc*>, const DeduceFunc*> KernelPack::GetKernel( |
| 28 | KernelPack::KernType kernel_type, Arch arch) { |
| 29 | //! arm64v7 is used by tinycv, nn opr should be armv64 or armv7, not arm64v7 |
| 30 | auto deduce_func = GetDeduceLayout(kernel_type); |
| 31 | if (arch == Arch::ARM64 || arch == Arch::ARM64V7 || arch == Arch::ARM64_WITH_I8MM) { |
| 32 | bool with_i8mm = (arch == Arch::ARM64_WITH_I8MM); |
| 33 | auto a64_kerns = Arm64::ArchKernelPack::GetKernel(kernel_type, with_i8mm); |
| 34 | auto armcommon_kerns = ArmCommon::ArchKernelPack::GetKernel(kernel_type); |
| 35 | auto gi_kerns = GeneralIntrinsic::ArchKernelPack::GetKernel(kernel_type); |
| 36 | if (kernel_type == KernelPack::KernType::MatrixMulKernel) { |
| 37 | armcommon_kerns.insert( |
| 38 | armcommon_kerns.end(), a64_kerns.begin(), a64_kerns.end()); |
| 39 | armcommon_kerns.insert( |
| 40 | armcommon_kerns.end(), gi_kerns.begin(), gi_kerns.end()); |
| 41 | return {armcommon_kerns, deduce_func}; |
| 42 | } |
| 43 | |
| 44 | std::vector<const KernelFunc*> valid_kern; |
| 45 | if (kernel_type == KernelPack::KernType::ConvKernel) { |
| 46 | std::vector<const KernelFunc*> sorted_kern; |
| 47 | for (auto&& kern : gi_kerns) { |
| 48 | auto kern_sym = kern->GetKernelSymbol(nullptr); |
| 49 | auto is_f63 = |
| 50 | std::regex_match(kern_sym, std::regex("^GI.*_winograd_f63.*")); |
| 51 | auto is_f43 = |
| 52 | std::regex_match(kern_sym, std::regex("^GI.*_winograd_f43.*")); |
| 53 | auto if_match = is_f63 || is_f43; |
| 54 | if (!if_match) { |
| 55 | valid_kern.push_back(kern); |
| 56 | } else { |
| 57 | if (is_f43) { |
| 58 | sorted_kern.insert(sorted_kern.begin(), kern); |
| 59 | } else { |
| 60 | sorted_kern.insert(sorted_kern.end(), kern); |
| 61 | } |
| 62 | } |
| 63 | } |
| 64 | //! WARNING: the f63 and f43 must exist in GI kernel |
| 65 | if (arch == Arch::ARM64) { |
| 66 | a64_kerns.insert( |
| 67 | a64_kerns.begin(), sorted_kern.begin(), sorted_kern.end()); |
| 68 | } |
| 69 | } else { |
| 70 | valid_kern = gi_kerns; |
| 71 | } |
| 72 | |
| 73 | a64_kerns.insert( |
| 74 | a64_kerns.end(), armcommon_kerns.begin(), armcommon_kerns.end()); |
| 75 | a64_kerns.insert(a64_kerns.end(), valid_kern.begin(), valid_kern.end()); |
| 76 | return {a64_kerns, deduce_func}; |
| 77 | |
| 78 | } else if (arch == Arch::ARMV7 || arch == Arch::ARMV7_WITH_DOT) { |
| 79 | bool with_dot = arch == Arch::ARMV7_WITH_DOT; |
| 80 | auto a32_kerns = Armv7::ArchKernelPack::GetKernel(kernel_type, with_dot); |
| 81 | |
| 82 | auto armcommon_kerns = ArmCommon::ArchKernelPack::GetKernel(kernel_type); |
| 83 | auto gi_kerns = GeneralIntrinsic::ArchKernelPack::GetKernel(kernel_type); |
| 84 | a32_kerns.insert( |
nothing calls this directly
no test coverage detected