/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
with_cpu
tests/node_tests.cpp
173 строки
9 KB
kolkir
Add autograd disable and enable functionality
28 янв 2025, 00:19
28 янв 2025, 00:19
92c91bd
Код
Авторство
О чём код?
#include <adept/autograd/autograd.hpp> #include <adept/nn/cross_entropy.hpp> #include <adept/tensor.hpp> #include "catch.hpp" using namespace adept; TEST_CASE("Add scalars gradient", "[gradient]") { auto x = Variable(Tensor::from_values({1.f}, {1, 1}, device_t::CPU)); auto y = Variable(Tensor::from_values({2.f}, {1, 1}, device_t::CPU)); auto z = x + y; z.backward(); REQUIRE_THAT(x.grad().at<float32_t>({0, 0}), Catch::WithinRel(1.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 0}), Catch::WithinRel(1.0f, 0.001f)); } TEST_CASE("Fail multidim gradient", "[gradient]") { auto x = Variable(Tensor::from_values({1.f, 3.f}, {1, 2}, device_t::CPU)); auto y = Variable(Tensor::from_values({2.f, 4.f}, {1, 2}, device_t::CPU)); auto z = x + y; REQUIRE_THROWS(z.backward()); } TEST_CASE("Tensor sum gradient", "[gradient]") { auto x = Variable(Tensor::from_values({1.f, 3.f}, {1, 2}, device_t::CPU)); auto y = Variable(Tensor::from_values({2.f, 4.f}, {1, 2}, device_t::CPU)); auto z = x + y; sum(z).backward(); REQUIRE_THAT(x.grad().at<float32_t>({0, 0}), Catch::WithinRel(1.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 0}), Catch::WithinRel(1.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({0, 1}), Catch::WithinRel(1.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 1}), Catch::WithinRel(1.0f, 0.001f)); } TEST_CASE("Tensor sum(mean) gradient", "[gradient]") { auto x = Variable(Tensor::from_values({1.f, 3.f}, {1, 2}, device_t::CPU)); auto y = Variable(Tensor::from_values({2.f, 4.f}, {1, 2}, device_t::CPU)); auto z = x + y; mean(z).backward(); REQUIRE_THAT(x.grad().at<float32_t>({0, 0}), Catch::WithinRel(0.5f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 0}), Catch::WithinRel(0.5f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({0, 1}), Catch::WithinRel(0.5f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 1}), Catch::WithinRel(0.5f, 0.001f)); } TEST_CASE("Tensor sub gradient", "[gradient]") { auto x = Variable(Tensor::from_values({1.f, 3.f}, {1, 2}, device_t::CPU)); auto y = Variable(Tensor::from_values({2.f, 4.f}, {1, 2}, device_t::CPU)); auto z = x - y; sum(z).backward(); REQUIRE_THAT(x.grad().at<float32_t>({0, 0}), Catch::WithinRel(-1.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 0}), Catch::WithinRel(-1.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({0, 1}), Catch::WithinRel(-1.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 1}), Catch::WithinRel(-1.0f, 0.001f)); } TEST_CASE("Tensor mul gradient", "[gradient]") { auto x = Variable(Tensor::from_values({2.f, 3.f}, {1, 2}, device_t::CPU)); auto y = Variable(Tensor::from_values({2.f, 1.5f}, {1, 2}, device_t::CPU)); auto z = x * y; sum(z).backward(); REQUIRE_THAT(x.grad().at<float32_t>({0, 0}), Catch::WithinRel(2.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({0, 1}), Catch::WithinRel(1.5f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 0}), Catch::WithinRel(2.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 1}), Catch::WithinRel(3.0f, 0.001f)); } TEST_CASE("Tensor div gradient", "[gradient]") { auto x = Variable(Tensor::from_values({2.f, 3.f}, {1, 2}, device_t::CPU)); auto y = Variable(Tensor::from_values({2.f, 1.5f}, {1, 2}, device_t::CPU)); auto z = x / y; sum(z).backward(); REQUIRE_THAT(x.grad().at<float32_t>({0, 0}), Catch::WithinRel(0.5f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({0, 1}), Catch::WithinRel(0.6667f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 0}), Catch::WithinRel(-0.5f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 1}), Catch::WithinRel(-1.3333f, 0.001f)); } TEST_CASE("Tensor dot gradient", "[gradient]") { auto x = Variable( Tensor::from_values({1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f}, {2, 4}, device_t::CPU)); auto y = Variable( Tensor::from_values({7.f, 8.f, 1.f, 3.f, 9.f, 4.f, 5.f, 2.f}, {4, 2}, device_t::CPU)); auto z = matmul(x, y); sum(z).backward(); REQUIRE_THAT(x.grad().at<float32_t>({0, 0}), Catch::WithinRel(15.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({0, 1}), Catch::WithinRel(4.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({0, 2}), Catch::WithinRel(13.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({0, 3}), Catch::WithinRel(7.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({1, 0}), Catch::WithinRel(15.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({1, 1}), Catch::WithinRel(4.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({1, 2}), Catch::WithinRel(13.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({1, 3}), Catch::WithinRel(7.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 0}), Catch::WithinRel(6.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({0, 1}), Catch::WithinRel(6.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({1, 0}), Catch::WithinRel(8.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({1, 1}), Catch::WithinRel(8.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({2, 0}), Catch::WithinRel(10.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({2, 1}), Catch::WithinRel(10.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({3, 0}), Catch::WithinRel(12.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({3, 1}), Catch::WithinRel(12.0f, 0.001f)); } TEST_CASE("Tensor batch dot gradient", "[gradient]") { auto x = Variable(Tensor::from_values( {1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f}, {2, 2, 4}, device_t::CPU)); auto y = Variable(Tensor::from_values( {7.f, 8.f, 1.f, 3.f, 9.f, 4.f, 5.f, 2.f, 7.f, 8.f, 1.f, 3.f, 9.f, 4.f, 5.f, 2.f}, {2, 4, 2}, device_t::CPU)); auto z = matmul(x, y); sum(z).backward(); for (index_t b = 0; b < 2; ++b) { REQUIRE_THAT(x.grad().at<float32_t>({b, 0, 0}), Catch::WithinRel(15.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 0, 1}), Catch::WithinRel(4.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 0, 2}), Catch::WithinRel(13.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 0, 3}), Catch::WithinRel(7.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 1, 0}), Catch::WithinRel(15.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 1, 1}), Catch::WithinRel(4.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 1, 2}), Catch::WithinRel(13.0f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 1, 3}), Catch::WithinRel(7.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({b, 0, 0}), Catch::WithinRel(6.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({b, 0, 1}), Catch::WithinRel(6.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({b, 1, 0}), Catch::WithinRel(8.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({b, 1, 1}), Catch::WithinRel(8.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({b, 2, 0}), Catch::WithinRel(10.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({b, 2, 1}), Catch::WithinRel(10.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({b, 3, 0}), Catch::WithinRel(12.0f, 0.001f)); REQUIRE_THAT(y.grad().at<float32_t>({b, 3, 1}), Catch::WithinRel(12.0f, 0.001f)); } } TEST_CASE("Tensor log_softmax gradient", "[gradient]") { auto x = Variable(Tensor::from_values( {1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f}, {2, 2, 4}, device_t::CPU)); auto z = log_softmax_last_dim(x); for (index_t b = 0; b < 2; ++b) { REQUIRE_THAT(z.data().at<float32_t>({b, 0, 0}), Catch::WithinRel(-3.4402f, 0.001f)); REQUIRE_THAT(z.data().at<float32_t>({b, 0, 1}), Catch::WithinRel(-2.4402f, 0.001f)); REQUIRE_THAT(z.data().at<float32_t>({b, 0, 2}), Catch::WithinRel(-1.4402f, 0.001f)); REQUIRE_THAT(z.data().at<float32_t>({b, 0, 3}), Catch::WithinRel(-0.4402f, 0.001f)); REQUIRE_THAT(z.data().at<float32_t>({b, 1, 0}), Catch::WithinRel(-3.4402f, 0.001f)); REQUIRE_THAT(z.data().at<float32_t>({b, 1, 1}), Catch::WithinRel(-2.4402f, 0.001f)); REQUIRE_THAT(z.data().at<float32_t>({b, 1, 2}), Catch::WithinRel(-1.4402f, 0.001f)); REQUIRE_THAT(z.data().at<float32_t>({b, 1, 3}), Catch::WithinRel(-0.4402f, 0.001f)); } mean(z).backward(); for (index_t b = 0; b < 2; ++b) { REQUIRE_THAT(x.grad().at<float32_t>({b, 0, 0}), Catch::WithinRel(0.0545f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 0, 1}), Catch::WithinRel(0.0407f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 0, 2}), Catch::WithinRel(0.0033f, 0.01f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 0, 3}), Catch::WithinRel(-0.0985f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 1, 0}), Catch::WithinRel(0.0545f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 1, 1}), Catch::WithinRel(0.0407f, 0.001f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 1, 2}), Catch::WithinRel(0.0033f, 0.01f)); REQUIRE_THAT(x.grad().at<float32_t>({b, 1, 3}), Catch::WithinRel(-0.0985f, 0.001f)); } }