| 87 | } |
| 88 | |
| 89 | static void ScaleCPU(DataType kernel_dtype, |
| 90 | const phi::CPUContext& dev_ctx, |
| 91 | const phi::DenseTensor& x, |
| 92 | const Scalar& scale, |
| 93 | const Scalar& bias, |
| 94 | bool bias_after_scale, |
| 95 | phi::DenseTensor* dense_out) { |
| 96 | switch (kernel_dtype) { |
| 97 | case phi::DataType::FLOAT64: { |
| 98 | phi::ScaleKernel<double>( |
| 99 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 100 | break; |
| 101 | } |
| 102 | case phi::DataType::FLOAT32: { |
| 103 | phi::ScaleKernel<float>( |
| 104 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 105 | break; |
| 106 | } |
| 107 | case phi::DataType::BFLOAT16: { |
| 108 | phi::ScaleKernel<phi::dtype::bfloat16>( |
| 109 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 110 | break; |
| 111 | } |
| 112 | case phi::DataType::INT64: { |
| 113 | phi::ScaleKernel<int64_t>( |
| 114 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 115 | break; |
| 116 | } |
| 117 | case phi::DataType::INT32: { |
| 118 | phi::ScaleKernel<int32_t>( |
| 119 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 120 | break; |
| 121 | } |
| 122 | case phi::DataType::INT16: { |
| 123 | phi::ScaleKernel<int16_t>( |
| 124 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 125 | break; |
| 126 | } |
| 127 | case phi::DataType::INT8: { |
| 128 | phi::ScaleKernel<int8_t>( |
| 129 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 130 | break; |
| 131 | } |
| 132 | case phi::DataType::UINT8: { |
| 133 | phi::ScaleKernel<uint8_t>( |
| 134 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 135 | break; |
| 136 | } |
| 137 | default: { |
| 138 | PADDLE_THROW(common::errors::Fatal( |
| 139 | "Detected unsupported data type." |
| 140 | "Only Float64, Float32, BFloat16, Int64, Int32, Int16, Int8, UInt8 " |
| 141 | "are supported for now.")); |
| 142 | break; |
| 143 | } |
| 144 | } |
| 145 | } |
| 146 | |