/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
tests/adam_tests.cpp
70 строк
4 KB
kolkir
Adam optimizer implementation
01 мар 2025, 17:03
01 мар 2025, 17:03
2fcc751
Код
Авторство
О чём код?
#include <adept/nn/adam.hpp> #include "catch.hpp" using namespace adept; TEST_CASE("Adam step in loop", "[adam]") { std::vector<Variable> params{Variable(Tensor::from_values({1.f, 2.f}, {2}, device_t::CPU)), Variable(Tensor::from_values({3.f, 4.f}, {2}, device_t::CPU))}; Adam optim(params); index_t n = 4; for (index_t i = 0; i < n; ++i) { for (auto& param : params) { param.grad() = Tensor::from_values({0.03f, 0.02f}, {2}, device_t::CPU); } optim.step(); if (i == 0) { auto state = optim.state(params[0].data().impl().get()); REQUIRE_THAT(state->exp_avg.at<float32_t>({0}), Catch::WithinRel(0.003f, 0.0001f)); REQUIRE_THAT(state->exp_avg.at<float32_t>({1}), Catch::WithinRel(0.002f, 0.0001f)); state = optim.state(params[1].data().impl().get()); REQUIRE_THAT(state->exp_avg.at<float32_t>({0}), Catch::WithinRel(0.003f, 0.0001f)); REQUIRE_THAT(state->exp_avg.at<float32_t>({1}), Catch::WithinRel(0.002f, 0.0001f)); REQUIRE_THAT(params[0].data().at<float32_t>({0}), Catch::WithinRel(0.999f, 0.0001f)); REQUIRE_THAT(params[0].data().at<float32_t>({1}), Catch::WithinRel(1.999f, 0.0001f)); REQUIRE_THAT(params[1].data().at<float32_t>({0}), Catch::WithinRel(2.999f, 0.0001f)); REQUIRE_THAT(params[1].data().at<float32_t>({1}), Catch::WithinRel(3.999f, 0.0001f)); } else if (i == 1) { auto state = optim.state(params[0].data().impl().get()); REQUIRE_THAT(state->exp_avg.at<float32_t>({0}), Catch::WithinRel(0.0057f, 0.0001f)); REQUIRE_THAT(state->exp_avg.at<float32_t>({1}), Catch::WithinRel(0.0038f, 0.0001f)); state = optim.state(params[1].data().impl().get()); REQUIRE_THAT(state->exp_avg.at<float32_t>({0}), Catch::WithinRel(0.0057f, 0.0001f)); REQUIRE_THAT(state->exp_avg.at<float32_t>({1}), Catch::WithinRel(0.0038f, 0.0001f)); REQUIRE_THAT(params[0].data().at<float32_t>({0}), Catch::WithinRel(0.998f, 0.0001f)); REQUIRE_THAT(params[0].data().at<float32_t>({1}), Catch::WithinRel(1.998f, 0.0001f)); REQUIRE_THAT(params[1].data().at<float32_t>({0}), Catch::WithinRel(2.998f, 0.0001f)); REQUIRE_THAT(params[1].data().at<float32_t>({1}), Catch::WithinRel(3.998f, 0.0001f)); } else if (i == 2) { auto state = optim.state(params[0].data().impl().get()); REQUIRE_THAT(state->exp_avg.at<float32_t>({0}), Catch::WithinRel(0.0081f, 0.01f)); REQUIRE_THAT(state->exp_avg.at<float32_t>({1}), Catch::WithinRel(0.0054f, 0.01f)); state = optim.state(params[1].data().impl().get()); REQUIRE_THAT(state->exp_avg.at<float32_t>({0}), Catch::WithinRel(0.0081f, 0.01f)); REQUIRE_THAT(state->exp_avg.at<float32_t>({1}), Catch::WithinRel(0.0054f, 0.01f)); REQUIRE_THAT(params[0].data().at<float32_t>({0}), Catch::WithinRel(0.997f, 0.0001f)); REQUIRE_THAT(params[0].data().at<float32_t>({1}), Catch::WithinRel(1.997f, 0.0001f)); REQUIRE_THAT(params[1].data().at<float32_t>({0}), Catch::WithinRel(2.997f, 0.0001f)); REQUIRE_THAT(params[1].data().at<float32_t>({1}), Catch::WithinRel(3.997f, 0.0001f)); } else if (i == 3) { auto state = optim.state(params[0].data().impl().get()); REQUIRE_THAT(state->exp_avg.at<float32_t>({0}), Catch::WithinRel(0.0103f, 0.01f)); REQUIRE_THAT(state->exp_avg.at<float32_t>({1}), Catch::WithinRel(0.0069f, 0.01f)); state = optim.state(params[1].data().impl().get()); REQUIRE_THAT(state->exp_avg.at<float32_t>({0}), Catch::WithinRel(0.0103f, 0.01f)); REQUIRE_THAT(state->exp_avg.at<float32_t>({1}), Catch::WithinRel(0.0069f, 0.01f)); REQUIRE_THAT(params[0].data().at<float32_t>({0}), Catch::WithinRel(0.996f, 0.0001f)); REQUIRE_THAT(params[0].data().at<float32_t>({1}), Catch::WithinRel(1.996f, 0.0001f)); REQUIRE_THAT(params[1].data().at<float32_t>({0}), Catch::WithinRel(2.996f, 0.0001f)); REQUIRE_THAT(params[1].data().at<float32_t>({1}), Catch::WithinRel(3.996f, 0.0001f)); } } }