/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/backends/cpu/batchnorm2d.cpp
326 строк
13 KB
kolkir
Fix small vector reduce
08 мар 2025, 15:05
08 мар 2025, 15:05
1870102
Код
Авторство
О чём код?
#undef HWY_TARGET_INCLUDE #define HWY_TARGET_INCLUDE "../src/backends/cpu/batchnorm2d.cpp" #include <hwy/foreach_target.h> #include <hwy/highway.h> #include <hwy/per_target.h> #include "batchnorm2d-inl.hpp" #if HWY_ONCE #include <adept/backends/cpu/batchnorm2d.hpp> #include <adept/irange.hpp> #include <adept/threading.hpp> #include <adept/types_dispatch.hpp> #include "data_index.hpp" namespace adept::cpu { namespace { template <typename T> T inv_std(T var, T epsilon) { T invstd = 0; if (var != static_cast<T>(0) || epsilon != static_cast<T>(0)) { invstd = static_cast<T>(1) / std::sqrt(var + epsilon); } return invstd; } template <typename scalar_t> void collect_stats(const Tensor& input, Tensor& mean, Tensor& var_sum) { auto n_batch = input.shape().dim(0); auto n_channel = input.shape().dim(1); auto channel_size = input.shape().numel() / n_batch / n_channel; auto n_per_channel = input.shape().numel() / n_channel; const auto* input_data = input.const_data_ptr<scalar_t>(); auto* mean_data = mean.mutable_data_ptr<scalar_t>(); auto* var_sum_data = var_sum.mutable_data_ptr<scalar_t>(); // parallel dim reduce on 'channel' parallel_for<scalar_t>(0, n_channel, [&](auto begin, auto end) { for (const auto c : irange(begin, end)) { // compute mean per input float64_t sum = 0; for (const auto n : irange(n_batch)) { for (const auto i : irange(channel_size)) { auto offset = n * n_channel * channel_size + c * channel_size + i; sum += input_data[offset]; } } scalar_t mean = sum / n_per_channel; mean_data[c] = mean; // compute variance per input float64_t var_sum = 0; for (const auto n : irange(n_batch)) { for (const auto i : irange(channel_size)) { auto offset = n * n_channel * channel_size + c * channel_size + i; auto x = input_data[offset]; var_sum += (x - mean) * (x - mean); } } var_sum_data[c] = var_sum; } }); } template <typename scalar_t> void update_stats(const Tensor& input, Tensor& running_mean, Tensor& running_var, float32_t momentum, float32_t eps, Tensor& save_mean, Tensor& save_var) { auto n_input = input.shape().dim(1); CHECK(input.shape().numel() != 0, "input tensor must have at least one element"); auto n = input.shape().numel() / n_input; auto mean = Tensor::empty({.shape = {n_input}, .device = input.device(), .dtype = dtype_t::Float32}); auto var_sum = Tensor::empty({.shape = {n_input}, .device = input.device(), .dtype = dtype_t::Float32}); auto mean_data_ptr = mean.mutable_data_ptr<scalar_t>(); auto var_sum_data_ptr = var_sum.mutable_data_ptr<scalar_t>(); collect_stats<scalar_t>(input, mean, var_sum); auto save_mean_data_ptr = save_mean.mutable_data_ptr<scalar_t>(); auto save_var_data_ptr = save_var.mutable_data_ptr<scalar_t>(); auto running_mean_data_ptr = running_mean.mutable_data_ptr<scalar_t>(); auto running_var_data_ptr = running_var.mutable_data_ptr<scalar_t>(); parallel_for<scalar_t>(0, n_input, [&](auto b_begin, auto b_end) { for (const auto f : irange(b_begin, b_end)) { save_mean_data_ptr[f] = mean_data_ptr[f]; save_var_data_ptr[f] = inv_std<scalar_t>(var_sum_data_ptr[f] / n, eps); running_mean_data_ptr[f] = momentum * mean_data_ptr[f] + (1 - momentum) * running_mean_data_ptr[f]; auto unbiased_var = var_sum_data_ptr[f] / (n - 1); running_var_data_ptr[f] = momentum * unbiased_var + (1 - momentum) * running_var_data_ptr[f]; } }); } template <typename scalar_t> void collect_linear_and_constant_terms(scalar_t* alpha, scalar_t* beta, int64_t n_channel, const Tensor& weight, const Tensor& bias, const Tensor& save_mean, const Tensor& save_invstd, const Tensor& running_mean, const Tensor& running_var, bool train, double eps) { const auto* weight_data = weight.const_data_ptr<scalar_t>(); const auto* bias_data = bias.const_data_ptr<scalar_t>(); const auto* save_mean_ptr = save_mean.const_data_ptr<scalar_t>(); const auto* save_invstd_ptr = save_invstd.const_data_ptr<scalar_t>(); const auto* running_mean_ptr = running_mean.const_data_ptr<scalar_t>(); const auto* running_var_ptr = running_var.const_data_ptr<scalar_t>(); for (const auto c : irange(n_channel)) { scalar_t mean, invstd; if (train) { mean = save_mean_ptr[c]; invstd = save_invstd_ptr[c]; } else { mean = running_mean_ptr[c]; invstd = 1 / std::sqrt(running_var_ptr[c] + static_cast<scalar_t>(eps)); } auto weight_v = weight_data[c]; auto bias_v = bias_data[c]; alpha[c] = invstd * weight_v; beta[c] = bias_v - mean * alpha[c]; } } template <typename scalar_t> void transform_input(const Tensor& input, const Tensor& weight, const Tensor& bias, Tensor& save_mean, Tensor& save_var, Tensor& running_mean, Tensor& running_var, Tensor& output, float32_t eps, bool train) { auto n_batch = input.shape().dim(0); auto n_channel = input.shape().dim(1); auto channel_size = input.shape().numel() / n_batch / n_channel; auto alpha = Tensor::empty({.shape = {n_channel}, .device = input.device(), .dtype = input.dtype()}); auto beta = Tensor::empty({.shape = {n_channel}, .device = input.device(), .dtype = input.dtype()}); auto* alpha_data = alpha.mutable_data_ptr<scalar_t>(); auto* beta_data = beta.mutable_data_ptr<scalar_t>(); collect_linear_and_constant_terms<scalar_t>(alpha_data, beta_data, n_channel, weight, bias, save_mean, save_var, running_mean, running_var, train, eps); auto* output_data = output.mutable_data_ptr<scalar_t>(); const auto* input_data = input.const_data_ptr<scalar_t>(); parallel_for<scalar_t>(0, n_batch * n_channel, [&](auto begin, auto end) { index_t n = 0; index_t c = 0; data_index_init(begin, n, n_batch, c, n_channel); for (const auto i : irange(begin, end)) { auto alpha_value = alpha_data[c]; auto beta_value = beta_data[c]; auto offset = i * channel_size; HWY_EXPORT_AND_DYNAMIC_DISPATCH_T(simd_batchnorm_fwd<scalar_t>) (output_data + offset, input_data + offset, alpha_value, beta_value, channel_size, input.is_shape_aligned()); // move on to next index data_index_step(n, n_batch, c, n_channel); } }); } template <typename scalar_t> void batchnorm2d_backward(Tensor& in_grad, Tensor& weight_grad, Tensor& bias_grad, const Tensor& out_grad, const Tensor& input, const Tensor& weight, const Tensor& save_mean, const Tensor& save_var, float32_t eps) { auto n_batch = input.shape().dim(0); auto n_channel = input.shape().dim(1); auto channel_size = input.shape().numel() / n_batch / n_channel; auto N = input.shape().numel() / n_channel; const auto* out_grad_data = out_grad.const_data_ptr<scalar_t>(); const auto* input_data = input.const_data_ptr<scalar_t>(); auto* in_grad_data = in_grad.defined() ? in_grad.mutable_data_ptr<scalar_t>() : nullptr; auto* weight_grad_data = weight_grad.mutable_data_ptr<scalar_t>(); auto* bias_grad_data = bias_grad.mutable_data_ptr<scalar_t>(); const bool in_grad_null = in_grad_data == nullptr; const auto* weight_ptr = weight.const_data_ptr<scalar_t>(); const auto* save_mean_ptr = save_mean.const_data_ptr<scalar_t>(); const auto* save_var_ptr = save_var.const_data_ptr<scalar_t>(); // parallel dim reduce on 'channel' parallel_for<scalar_t>(0, n_channel, [&](auto begin, auto end) { for (const auto c : irange(begin, end)) { scalar_t w = weight_ptr[c]; scalar_t mean = save_mean_ptr[c]; scalar_t var = save_var_ptr[c]; // reduce over grad_output in feature plane // compute sum and dot product of Q(X) and dY. // fuse into a single loop to reuse dY scalar_t sum = 0; scalar_t dotp = 0; for (const auto n : irange(n_batch)) { const auto* x_ptr = input_data + n * n_channel * channel_size + c * channel_size; const auto* dy_ptr = out_grad_data + n * n_channel * channel_size + c * channel_size; HWY_EXPORT_AND_DYNAMIC_DISPATCH_T(simd_bn_sum<scalar_t>) (sum, dy_ptr, channel_size, input.is_shape_aligned()); HWY_EXPORT_AND_DYNAMIC_DISPATCH_T(simd_bn_dotp<scalar_t>) (dotp, mean, x_ptr, dy_ptr, channel_size, input.is_shape_aligned()); } if (!in_grad_null) { auto k = dotp * var * var / N; auto grad_mean = sum / N; for (const auto n : irange(n_batch)) { const auto* x_ptr = input_data + n * n_channel * channel_size + c * channel_size; auto* dx_ptr = in_grad_data + n * n_channel * channel_size + c * channel_size; const auto* dy_ptr = out_grad_data + n * n_channel * channel_size + c * channel_size; HWY_EXPORT_AND_DYNAMIC_DISPATCH_T(simd_bn_in_grad<scalar_t>) (grad_mean, mean, k, var, w, dx_ptr, x_ptr, dy_ptr, channel_size, input.is_shape_aligned()); } } weight_grad_data[c] = dotp * var; bias_grad_data[c] = sum; } }); } } // namespace /* output(n, c, h, w) = (input(n, c, h, w) - mean(c)) / sqrt(var(c) + eps) * weight(c) + bias(c) = input(n, c, h, w) * inv_var(c) * weight(c) - mean(c) * inv_var(c) * weight(c) + bias(c) where inv_var(c) = 1 / sqrt(var(c) + eps) So the linear term: alpha(c) = inv_var(c) * weight(c), The constant term: beta(c) = bias(c) - mean(c) * inv_var(c) * weight(c) */ std::tuple<Tensor, Tensor, Tensor> batchnorm2d_fwd(const Tensor& input, const Tensor& weight, const Tensor& bias, Tensor& running_mean, Tensor& running_var, float32_t eps, float32_t momentum, bool is_train) { CHECK(input.shape().rang() == 4, "Invalid input shape ", input.shape(), " for batchnorm2d"); auto output = Tensor::zero({.shape = input.shape(), .device = input.device(), .dtype = input.dtype()}); auto save_mean = Tensor::zero( {.shape = {input.shape().dim(1)}, .device = input.device(), .dtype = input.dtype()}); auto save_var = Tensor::zero( {.shape = {input.shape().dim(1)}, .device = input.device(), .dtype = input.dtype()}); DISPATCH_FLOATING_TYPES(input.dtype(), [&]() { if (is_train) { update_stats<scalar_t>(input, running_mean, running_var, momentum, eps, save_mean, save_var); } transform_input<scalar_t>(input, weight, bias, save_mean, save_var, running_mean, running_var, output, eps, is_train); }); return {output, save_mean, save_var}; } std::tuple<Tensor, Tensor> batchnorm2d_bwd(Tensor& in_grad, const Tensor& out_grad, const Tensor& input, const Tensor& weight, const Tensor& save_mean, const Tensor& save_var, float32_t eps) { CHECK(input.shape().rang() == 4, "Invalid input shape ", input.shape(), " for batchnorm2d"); auto weigth_grad = Tensor::zero( {.shape = {input.shape().dim(1)}, .device = input.device(), .dtype = input.dtype()}); auto bias_grad = Tensor::zero( {.shape = {input.shape().dim(1)}, .device = input.device(), .dtype = input.dtype()}); DISPATCH_FLOATING_TYPES(input.dtype(), [&]() { batchnorm2d_backward<scalar_t>(in_grad, weigth_grad, bias_grad, out_grad, input, weight, save_mean, save_var, eps); }); return {weigth_grad, bias_grad}; } } // namespace adept::cpu #endif // HWY_ONCE