| 206 | #endif |
| 207 | |
| 208 | Tensor scale_switch_case(const Tensor& x, |
| 209 | const Scalar& scale, |
| 210 | const Scalar& bias, |
| 211 | bool bias_after_scale) { |
| 212 | Backend kernel_backend = Backend::UNDEFINED; |
| 213 | DataLayout kernel_layout = DataLayout::UNDEFINED; |
| 214 | DataType kernel_data_type = DataType::UNDEFINED; |
| 215 | |
| 216 | if (kernel_backend == Backend::UNDEFINED || |
| 217 | kernel_layout == DataLayout::UNDEFINED || |
| 218 | kernel_data_type == DataType::UNDEFINED) { |
| 219 | auto kernel_key_set = ParseKernelKeyByInputArgs(x); |
| 220 | auto kernel_key = kernel_key_set.GetHighestPriorityKernelKey(); |
| 221 | if (kernel_backend == Backend::UNDEFINED) { |
| 222 | kernel_backend = kernel_key.backend(); |
| 223 | } |
| 224 | if (kernel_layout == DataLayout::UNDEFINED) { |
| 225 | kernel_layout = kernel_key.layout(); |
| 226 | } |
| 227 | if (kernel_data_type == DataType::UNDEFINED) { |
| 228 | kernel_data_type = kernel_key.dtype(); |
| 229 | } |
| 230 | } |
| 231 | auto kernel_result = phi::KernelFactory::Instance().SelectKernelOrThrowError( |
| 232 | "scale", {kernel_backend, kernel_layout, kernel_data_type}); |
| 233 | const auto& kernel = kernel_result.kernel; |
| 234 | if (FLAGS_low_precision_op_list) { |
| 235 | phi::KernelFactory::Instance().AddToLowPrecisionKernelList( |
| 236 | "scale", kernel_data_type); |
| 237 | } |
| 238 | VLOG(6) << "scale API kernel key: [" << kernel_backend << ", " |
| 239 | << kernel_layout << ", " << kernel_data_type << "]"; |
| 240 | VLOG(6) << "scale API kernel: " << kernel; |
| 241 | |
| 242 | auto* dev_ctx = GetDeviceContextByBackend(kernel_backend); |
| 243 | |
| 244 | auto dense_x = std::dynamic_pointer_cast<phi::DenseTensor>(x.impl()); |
| 245 | |
| 246 | auto dense_out = std::make_shared<phi::DenseTensor>(); |
| 247 | phi::MetaTensor meta_out(dense_out.get()); |
| 248 | phi::UnchangedInferMeta(*dense_x, &meta_out); |
| 249 | |
| 250 | Tensor out; |
| 251 | out.set_impl(dense_out); |
| 252 | |
| 253 | switch (kernel_backend) { |
| 254 | case Backend::CPU: |
| 255 | ScaleCPU(kernel_data_type, |
| 256 | static_cast<const phi::CPUContext&>(*dev_ctx), |
| 257 | *dense_x, |
| 258 | scale, |
| 259 | bias, |
| 260 | bias_after_scale, |
| 261 | dense_out.get()); |
| 262 | break; |
| 263 | #if defined(PADDLE_WITH_CUDA) || defined(PADDLE_WITH_HIP) |
| 264 | case Backend::GPU: |
| 265 | ScaleGPU(kernel_data_type, |