| 265 | } |
| 266 | |
| 267 | KernelResult KernelFactory::SelectKernelOrThrowError( |
| 268 | const std::string& kernel_name, |
| 269 | const KernelKey& const_kernel_key, |
| 270 | bool use_strided_kernel) const { |
| 271 | auto iter = kernels_.find(kernel_name); |
| 272 | |
| 273 | PADDLE_ENFORCE_NE(iter, |
| 274 | kernels_.end(), |
| 275 | common::errors::NotFound( |
| 276 | "The kernel `%s` is not registered.", kernel_name)); |
| 277 | if (FLAGS_use_stride_kernel && use_strided_kernel) { |
| 278 | auto stride_kernel_iter = iter->second.find( |
| 279 | {const_kernel_key.backend() == paddle::experimental::Backend::GPUDNN |
| 280 | ? paddle::experimental::get_accelerat_backend() |
| 281 | : const_kernel_key.backend(), |
| 282 | phi::DataLayout::STRIDED, |
| 283 | const_kernel_key.dtype()}); |
| 284 | if (stride_kernel_iter != iter->second.end()) { |
| 285 | return {stride_kernel_iter->second, false, true}; |
| 286 | } |
| 287 | #ifdef PADDLE_WITH_CUSTOM_DEVICE |
| 288 | if (stride_kernel_iter == iter->second.end() && |
| 289 | (const_kernel_key.backend() > phi::Backend::NUM_BACKENDS || |
| 290 | const_kernel_key.backend() == phi::Backend::GPUDNN)) { |
| 291 | stride_kernel_iter = iter->second.find({phi::Backend::CUSTOM, |
| 292 | phi::DataLayout::STRIDED, |
| 293 | const_kernel_key.dtype()}); |
| 294 | if (stride_kernel_iter != iter->second.end()) { |
| 295 | return {stride_kernel_iter->second, false, true}; |
| 296 | } |
| 297 | } |
| 298 | #endif |
| 299 | } |
| 300 | |
| 301 | KernelKey kernel_key = KernelKey(const_kernel_key.backend(), |
| 302 | phi::DataLayout::ALL_LAYOUT, |
| 303 | const_kernel_key.dtype()); |
| 304 | #if defined(PADDLE_WITH_CUDA) || defined(PADDLE_WITH_HIP) || \ |
| 305 | defined(PADDLE_WITH_CUSTOM_DEVICE) |
| 306 | if (kernel_key.backend() == Backend::GPUDNN) { |
| 307 | auto kernel_iter = iter->second.find( |
| 308 | {Backend::GPUDNN, phi::DataLayout::ALL_LAYOUT, kernel_key.dtype()}); |
| 309 | if (kernel_iter != iter->second.end()) { |
| 310 | return {kernel_iter->second, false, false}; |
| 311 | } |
| 312 | kernel_key = KernelKey(paddle::experimental::get_accelerat_backend(), |
| 313 | kernel_key.layout(), |
| 314 | kernel_key.dtype()); |
| 315 | } |
| 316 | #endif |
| 317 | auto kernel_iter = iter->second.find(kernel_key); |
| 318 | |
| 319 | PADDLE_ENFORCE_NE( |
| 320 | kernel_iter == iter->second.end() && kernel_key.backend() == Backend::CPU, |
| 321 | true, |
| 322 | common::errors::NotFound( |
| 323 | "The kernel with key %s of kernel `%s` is not registered. %s", |
| 324 | kernel_key, |