| 146 | |
| 147 | #if defined(PADDLE_WITH_CUDA) || defined(PADDLE_WITH_HIP) |
| 148 | static void ScaleGPU(DataType kernel_dtype, |
| 149 | const phi::GPUContext& dev_ctx, |
| 150 | const phi::DenseTensor& x, |
| 151 | const Scalar& scale, |
| 152 | const Scalar& bias, |
| 153 | bool bias_after_scale, |
| 154 | phi::DenseTensor* dense_out) { |
| 155 | switch (kernel_dtype) { |
| 156 | case phi::DataType::FLOAT64: { |
| 157 | phi::ScaleKernel<double>( |
| 158 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 159 | break; |
| 160 | } |
| 161 | case phi::DataType::FLOAT32: { |
| 162 | phi::ScaleKernel<float>( |
| 163 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 164 | break; |
| 165 | } |
| 166 | case phi::DataType::FLOAT16: { |
| 167 | phi::ScaleKernel<phi::dtype::float16>( |
| 168 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 169 | break; |
| 170 | } |
| 171 | case phi::DataType::INT64: { |
| 172 | phi::ScaleKernel<int64_t>( |
| 173 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 174 | break; |
| 175 | } |
| 176 | case phi::DataType::INT32: { |
| 177 | phi::ScaleKernel<int32_t>( |
| 178 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 179 | break; |
| 180 | } |
| 181 | case phi::DataType::INT16: { |
| 182 | phi::ScaleKernel<int16_t>( |
| 183 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 184 | break; |
| 185 | } |
| 186 | case phi::DataType::INT8: { |
| 187 | phi::ScaleKernel<int8_t>( |
| 188 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 189 | break; |
| 190 | } |
| 191 | case phi::DataType::UINT8: { |
| 192 | phi::ScaleKernel<uint8_t>( |
| 193 | dev_ctx, x, scale, bias, bias_after_scale, dense_out); |
| 194 | break; |
| 195 | } |
| 196 | default: { |
| 197 | PADDLE_THROW(common::errors::Fatal( |
| 198 | "Detected unsupported data type." |
| 199 | "Only Float64, Float32, Float16, Int64, Int32, Int16, Int8, UInt8 " |
| 200 | "are " |
| 201 | "supported for now.")); |
| 202 | break; |
| 203 | } |
| 204 | } |
| 205 | } |