| 105 | } |
| 106 | |
| 107 | bool PluginRegistry::SetDefaultFactory(Platform::Id platform_id, |
| 108 | PluginKind plugin_kind, |
| 109 | PluginId plugin_id) { |
| 110 | if (!HasFactory(platform_id, plugin_kind, plugin_id)) { |
| 111 | port::StatusOr<Platform*> status = |
| 112 | MultiPlatformManager::PlatformWithId(platform_id); |
| 113 | string platform_name = "<unregistered platform>"; |
| 114 | if (status.ok()) { |
| 115 | platform_name = status.ValueOrDie()->Name(); |
| 116 | } |
| 117 | |
| 118 | LOG(ERROR) << "A factory must be registered for a platform before being " |
| 119 | << "set as default! " |
| 120 | << "Platform name: " << platform_name |
| 121 | << ", PluginKind: " << PluginKindString(plugin_kind) |
| 122 | << ", PluginId: " << plugin_id; |
| 123 | return false; |
| 124 | } |
| 125 | |
| 126 | switch (plugin_kind) { |
| 127 | case PluginKind::kBlas: |
| 128 | default_factories_[platform_id].blas = plugin_id; |
| 129 | break; |
| 130 | case PluginKind::kDnn: |
| 131 | default_factories_[platform_id].dnn = plugin_id; |
| 132 | break; |
| 133 | case PluginKind::kFft: |
| 134 | default_factories_[platform_id].fft = plugin_id; |
| 135 | break; |
| 136 | case PluginKind::kRng: |
| 137 | default_factories_[platform_id].rng = plugin_id; |
| 138 | break; |
| 139 | default: |
| 140 | LOG(ERROR) << "Invalid plugin kind specified: " |
| 141 | << static_cast<int>(plugin_kind); |
| 142 | return false; |
| 143 | } |
| 144 | |
| 145 | return true; |
| 146 | } |
| 147 | |
| 148 | bool PluginRegistry::HasFactory(const PluginFactories& factories, |
| 149 | PluginKind plugin_kind, |
no test coverage detected