/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
with_cpu
src/nn/adam.cpp
74 строки
2 KB
kolkir
Fix training issues
06 мар 2025, 23:54
06 мар 2025, 23:54
efe26e8
Код
Авторство
О чём код?
#include <adept/nn/adam.hpp> #include <adept/threading.hpp> #include <cmath> namespace adept { Adam::Adam(std::vector<Variable> parameters, float32_t lr, float32_t beta1, float32_t beta2, float32_t eps) : parameters_(parameters), lr_(lr), beta1_(beta1), beta2_(beta2), eps_(eps) {} void Adam::step() { // TODO: consider to use parallel_for? needs benchmarking. for (auto& param : parameters_) { if (param.defined() && param.requires_grad()) { auto grad = param.grad(); auto data = param.data(); auto param_state = state_.find(data.impl().get()); // initialize state if (param_state == state_.end()) { auto state = std::make_unique<detail::AdamParamState>(); state->step = 0; // exponential moving average of gradient values state->exp_avg = Tensor::zero(grad.properties()); // exponential moving average of squared gradient values state->exp_avg_sq = Tensor::zero(grad.properties()); state_[data.impl().get()] = std::move(state); } auto& state = static_cast<detail::AdamParamState&>(*state_[data.impl().get()]); auto& exp_avg = state.exp_avg; auto& exp_avg_sq = state.exp_avg_sq; state.step += 1; auto bias_correction1 = 1 - std::pow(beta1_, state.step); auto bias_correction2 = 1 - std::pow(beta2_, state.step); // decay the first and second moment running average coefficient exp_avg.mul_(beta1_).add_(grad * (1 - beta1_)); exp_avg_sq.mul_(beta2_).add_((grad * grad) * (1 - beta2_)); auto denom = (exp_avg_sq.sqrt() / std::sqrt(bias_correction2)).add_(eps_); auto step_size = lr_ / bias_correction1; data -= (exp_avg / denom) * step_size; } } } void Adam::zero_grad() { parallel_for_each(parameters_, [](auto& param) { if (param.defined()) param.zero_grad(); }); } void Adam::set_lr(float32_t lr) { lr_ = lr; } float32_t Adam::lr() const { return lr_; } const detail::AdamParamState* Adam::state(void* key) const { return state_.at(key).get(); } } // namespace adept