| 31 | namespace experimental { |
| 32 | |
| 33 | Tensor scale_kernel_context(const Tensor& x, |
| 34 | const Scalar& scale, |
| 35 | const Scalar& bias, |
| 36 | bool bias_after_scale) { |
| 37 | Backend kernel_backend = Backend::UNDEFINED; |
| 38 | DataLayout kernel_layout = DataLayout::UNDEFINED; |
| 39 | DataType kernel_data_type = DataType::UNDEFINED; |
| 40 | |
| 41 | if (kernel_backend == Backend::UNDEFINED || |
| 42 | kernel_layout == DataLayout::UNDEFINED || |
| 43 | kernel_data_type == DataType::UNDEFINED) { |
| 44 | auto kernel_key_set = ParseKernelKeyByInputArgs(x); |
| 45 | auto kernel_key = kernel_key_set.GetHighestPriorityKernelKey(); |
| 46 | if (kernel_backend == Backend::UNDEFINED) { |
| 47 | kernel_backend = kernel_key.backend(); |
| 48 | } |
| 49 | if (kernel_layout == DataLayout::UNDEFINED) { |
| 50 | kernel_layout = kernel_key.layout(); |
| 51 | } |
| 52 | if (kernel_data_type == DataType::UNDEFINED) { |
| 53 | kernel_data_type = kernel_key.dtype(); |
| 54 | } |
| 55 | } |
| 56 | auto kernel_result = phi::KernelFactory::Instance().SelectKernelOrThrowError( |
| 57 | "scale", {kernel_backend, kernel_layout, kernel_data_type}); |
| 58 | const auto& kernel = kernel_result.kernel; |
| 59 | if (FLAGS_low_precision_op_list) { |
| 60 | phi::KernelFactory::Instance().AddToLowPrecisionKernelList( |
| 61 | "scale", kernel_data_type); |
| 62 | } |
| 63 | VLOG(6) << "scale API kernel key: [" << kernel_backend << ", " |
| 64 | << kernel_layout << ", " << kernel_data_type << "]"; |
| 65 | VLOG(6) << "scale API kernel: " << kernel; |
| 66 | |
| 67 | auto* dev_ctx = GetDeviceContextByBackend(kernel_backend); |
| 68 | auto kernel_context = phi::KernelContext(dev_ctx); |
| 69 | |
| 70 | auto dense_x = std::dynamic_pointer_cast<phi::DenseTensor>(x.impl()); |
| 71 | kernel_context.EmplaceBackInput(dense_x.get()); |
| 72 | |
| 73 | kernel_context.EmplaceBackAttr(scale); |
| 74 | kernel_context.EmplaceBackAttr(bias); |
| 75 | kernel_context.EmplaceBackAttr(bias_after_scale); |
| 76 | |
| 77 | auto dense_out = std::make_shared<phi::DenseTensor>(); |
| 78 | phi::MetaTensor meta_out(dense_out.get()); |
| 79 | phi::UnchangedInferMeta(*dense_x, &meta_out); |
| 80 | kernel_context.EmplaceBackOutput(dense_out.get()); |
| 81 | |
| 82 | Tensor out; |
| 83 | out.set_impl(dense_out); |
| 84 | |
| 85 | kernel(&kernel_context); |
| 86 | return out; |
| 87 | } |
| 88 | |
| 89 | static void ScaleCPU(DataType kernel_dtype, |
| 90 | const phi::CPUContext& dev_ctx, |