MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / SelectKernelOrThrowError

Method SelectKernelOrThrowError

paddle/phi/core/kernel_factory.cc:267–643  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

265}
266
267KernelResult 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,

Callers 15

add_n_implFunction · 0.80
fused_gemm_epilogue_implFunction · 0.80
cudnn_lstm_grad_implFunction · 0.80
embedding_grad_implFunction · 0.80
TransDataTypeFunction · 0.80
PhiKernelInstructionMethod · 0.80
RunMethod · 0.80
CopyOrAddTensorFunction · 0.80
operator()Method · 0.80
operator()Method · 0.80

Calls 15

get_accelerat_backendFunction · 0.85
is_xpu_kp_support_opFunction · 0.85
is_xpu_support_opFunction · 0.85
is_in_custom_black_listFunction · 0.85
args_defMethod · 0.80
input_defsMethod · 0.80
output_defsMethod · 0.80
KernelKeyClass · 0.70
findMethod · 0.45
endMethod · 0.45
backendMethod · 0.45

Tested by 1

TESTFunction · 0.64