| 21 | |
| 22 | template <typename T, typename Context> |
| 23 | void InstanceNormKernel(const Context& dev_ctx, |
| 24 | const DenseTensor& x, |
| 25 | const optional<DenseTensor>& scale, |
| 26 | const optional<DenseTensor>& bias, |
| 27 | float epsilon, |
| 28 | DenseTensor* y, |
| 29 | DenseTensor* saved_mean, |
| 30 | DenseTensor* saved_var) { |
| 31 | using XPUType = typename XPUTypeTrait<T>::Type; |
| 32 | |
| 33 | const auto& x_dims = x.dims(); |
| 34 | int64_t n = x_dims[0]; |
| 35 | int64_t c = x_dims[1]; |
| 36 | int64_t h = x_dims[2]; |
| 37 | int64_t w = x_dims[3]; |
| 38 | dev_ctx.template Alloc<T>(y); |
| 39 | dev_ctx.template Alloc<float>(saved_mean); |
| 40 | dev_ctx.template Alloc<float>(saved_var); |
| 41 | if (x.numel() == 0) { |
| 42 | if (y) { |
| 43 | Full<T, Context>(dev_ctx, y->dims(), 0, y); |
| 44 | } |
| 45 | if (saved_mean) { |
| 46 | Full<float, Context>(dev_ctx, saved_mean->dims(), 0.f, saved_mean); |
| 47 | } |
| 48 | if (saved_var) { |
| 49 | Full<float, Context>(dev_ctx, saved_var->dims(), 0.f, saved_var); |
| 50 | } |
| 51 | return; |
| 52 | } |
| 53 | |
| 54 | xpu::ctx_guard RAII_GUARD(dev_ctx.x_context()); |
| 55 | |
| 56 | // scale |
| 57 | const auto scale_ptr = scale.get_ptr(); |
| 58 | const float* scale_data_fp32 = nullptr; |
| 59 | if (scale_ptr == nullptr) { |
| 60 | float* scale_data_temp = RAII_GUARD.alloc_l3_or_gm<float>(c); |
| 61 | int r = xpu::constant<float>(dev_ctx.x_context(), scale_data_temp, c, 1.f); |
| 62 | PADDLE_ENFORCE_XDNN_SUCCESS(r, "constant"); |
| 63 | scale_data_fp32 = scale_data_temp; |
| 64 | } else if (scale_ptr->dtype() == |
| 65 | phi::CppTypeToDataType<phi::float16>::Type()) { |
| 66 | float* scale_data_temp = |
| 67 | RAII_GUARD.alloc_l3_or_gm<float>(scale_ptr->numel()); |
| 68 | int r = xpu::cast<XPUType, float>( |
| 69 | dev_ctx.x_context(), |
| 70 | reinterpret_cast<const XPUType*>(scale_ptr->data<T>()), |
| 71 | scale_data_temp, |
| 72 | scale_ptr->numel()); |
| 73 | PADDLE_ENFORCE_XDNN_SUCCESS(r, "cast"); |
| 74 | scale_data_fp32 = scale_data_temp; |
| 75 | } else { |
| 76 | // no need to cast |
| 77 | scale_data_fp32 = scale_ptr->data<float>(); |
| 78 | } |
| 79 | |
| 80 | // bias |