/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/nn/batchnorm2d.cpp
92 строки
3 KB
kolkir
Make test convnet cpp app
07 мар 2025, 10:56
07 мар 2025, 10:56
997ae5b
Код
Авторство
О чём код?
#include <adept/nn/batchnorm2d.hpp> #include <adept/nn/init.hpp> #include "../dispatch/batchnorm.hpp" namespace adept { BarchNorm2dImpl::BarchNorm2dImpl(index_t num_features, float32_t eps, float32_t momentum, device_t device, dtype_t dtype) : num_features_(num_features), eps_(eps), momentum_(momentum), weight_(Tensor::empty({.shape = {num_features}, .device = device, .dtype = dtype})), bias_(Tensor::empty({.shape = {num_features}, .device = device, .dtype = dtype})), running_mean_(Tensor::empty({.shape = {num_features}, .device = device, .dtype = dtype})), running_var_(Tensor::empty({.shape = {num_features}, .device = device, .dtype = dtype})) { init_weights(); set_name("BatchNorm2d"); register_parameter("weight", weight_); register_parameter("bias", bias_); // TODO: consider add buffers registering for load/store automating running_mean_.requires_grad(false); register_parameter("running_mean_", running_mean_); running_var_.requires_grad(false); register_parameter("running_var", running_var_); } void BarchNorm2dImpl::init_weights() { fill(weight_.data(), 1.f); fill_zero(bias_.data()); fill_zero(running_mean_.data()); fill_zero(running_var_.data()); } void BarchNorm2dImpl::set_weights(const Tensor& weight) { weight_.data() = weight; } void BarchNorm2dImpl::set_bias(const Tensor& bias) { bias_.data() = bias; } Variable BarchNorm2dImpl::weights() const { return weight_; } Variable BarchNorm2dImpl::bias() const { return bias_; } void BarchNorm2dImpl::set_running_mean(const Tensor& mean) { running_mean_.data() = mean; } void BarchNorm2dImpl::set_running_var(const Tensor& var) { running_var_.data() = var; } Variable BarchNorm2dImpl::running_mean() const { return running_mean_; } Variable BarchNorm2dImpl::running_var() const { return running_var_; } Variable BarchNorm2dImpl::forward(const Variable& input) { CHECK(input.data().shape().dim(1) == num_features_, " BarchNorm2d input has invalid feature number ", num_features_, " != ", input.data().shape().dim(1)); auto [result, save_mean, save_var] = batchnorm2d_fwd(input.data(), weight_.data(), bias_.data(), running_mean_.data(), running_var_.data(), eps_, momentum_, is_train_mode_); Variable var(std::move(result), {input}, "batchnorm2d"); var.set_backward_fn([this, input = input, save_mean = save_mean, save_var = save_var](const auto& out_grad) mutable { Tensor input_grad; if (input.requires_grad()) input_grad = input.grad(); auto [weight_grad, bias_grad] = batchnorm2d_bwd(input_grad, out_grad, input.data(), weight_.data(), save_mean, save_var, eps_); weight_.add_grad(weight_grad); bias_.add_grad(bias_grad); }); return var; } } // namespace adept