/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
tests/conv2d_tests.cpp
151 строка
7 KB
kolkir
Conv2d backward pass implementation
02 фев 2025, 21:49
02 фев 2025, 21:49
9fe8c9a
Код
Авторство
О чём код?
#include <adept/nn/conv2d.hpp> #include "catch.hpp" using namespace adept; TEST_CASE("Conv2dOptions default construction", "[conv2d]") { Conv2dOptions opt(1, 1, 3); REQUIRE(opt.in_channels == 1); REQUIRE(opt.out_channels == 1); REQUIRE(opt.bias == true); REQUIRE(opt.kernel == ParamArray<2>{3, 3}); REQUIRE(opt.stride == ParamArray<2>{1, 1}); REQUIRE(opt.padding == ParamArray<2>{0, 0}); REQUIRE(opt.dilation == ParamArray<2>{1, 1}); } TEST_CASE("Conv2dOptions extended construction", "[conv2d]") { auto opt = Conv2dOptions(1, 1, 3) .with_padding({3, 1}) .with_bias(false) .with_stride({3, 1}) .with_dilation({3, 1}); REQUIRE(opt.in_channels == 1); REQUIRE(opt.out_channels == 1); REQUIRE(opt.bias == false); REQUIRE(opt.kernel == ParamArray<2>{3, 3}); REQUIRE(opt.stride == ParamArray<2>{3, 1}); REQUIRE(opt.padding == ParamArray<2>{3, 1}); REQUIRE(opt.dilation == ParamArray<2>{3, 1}); } TEST_CASE("Conv2dO construction", "[conv2d]") { Conv2d conv(Conv2dOptions(1, 1, 3)); } TEST_CASE("Conv2d im3x3 k2x2 ch1x1", "[conv2d]") { std::vector<float32_t> im_data = {1, 2, 3, 4, 5, 6, 7, 8, 9}; std::vector<float32_t> kernel_data = {7, 8, 1, 3}; auto im = Tensor::from_blob( im_data.data(), {.shape = {1, 1, 3, 3}, .device = device_t::CPU, .dtype = to_dtype<float32_t>()}); auto kernels = Tensor::from_blob( kernel_data.data(), {.shape = {1, 1, 2, 2}, .device = device_t::CPU, .dtype = to_dtype<float32_t>()}); Conv2d conv(Conv2dOptions(1, 1, 2).with_bias(false)); conv->set_weights(kernels); auto res = conv(im); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 0, 0}), Catch::WithinRel(42.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 0, 1}), Catch::WithinRel(61.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 1, 0}), Catch::WithinRel(99.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 1, 1}), Catch::WithinRel(118.f, 0.001f)); mean(res).backward(); REQUIRE_THAT(conv->weights().grad().at<float32_t>({0, 0, 0, 0}), Catch::WithinRel(3.f, 0.001f)); REQUIRE_THAT(conv->weights().grad().at<float32_t>({0, 0, 0, 1}), Catch::WithinRel(4.f, 0.001f)); REQUIRE_THAT(conv->weights().grad().at<float32_t>({0, 0, 1, 0}), Catch::WithinRel(6.f, 0.001f)); REQUIRE_THAT(conv->weights().grad().at<float32_t>({0, 0, 1, 1}), Catch::WithinRel(7.f, 0.001f)); } TEST_CASE("Conv2d im3x3 k2x2 ch1x2", "[conv2d]") { std::vector<float32_t> im_data = {1, 2, 3, 4, 5, 6, 7, 8, 9}; std::vector<float32_t> kernel_data = {7, 8, 1, 3, 7, 8, 1, 3}; auto im = Tensor::from_blob( im_data.data(), {.shape = {1, 1, 3, 3}, .device = device_t::CPU, .dtype = to_dtype<float32_t>()}); auto kernels = Tensor::from_blob( kernel_data.data(), {.shape = {2, 1, 2, 2}, .device = device_t::CPU, .dtype = to_dtype<float32_t>()}); Conv2d conv(Conv2dOptions(1, 2, 2).with_bias(false)); conv->set_weights(kernels); auto res = conv(im); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 0, 0}), Catch::WithinRel(42.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 0, 1}), Catch::WithinRel(61.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 1, 0}), Catch::WithinRel(99.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 1, 1}), Catch::WithinRel(118.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 1, 0, 0}), Catch::WithinRel(42.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 1, 0, 1}), Catch::WithinRel(61.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 1, 1, 0}), Catch::WithinRel(99.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 1, 1, 1}), Catch::WithinRel(118.f, 0.001f)); } TEST_CASE("Conv2d im3x3 k2x2 ch1x1 bias", "[conv2d]") { std::vector<float32_t> im_data = {1, 2, 3, 4, 5, 6, 7, 8, 9}; std::vector<float32_t> kernel_data = {7, 8, 1, 3}; auto im = Tensor::from_blob( im_data.data(), {.shape = {1, 1, 3, 3}, .device = device_t::CPU, .dtype = to_dtype<float32_t>()}); auto kernels = Tensor::from_blob( kernel_data.data(), {.shape = {1, 1, 2, 2}, .device = device_t::CPU, .dtype = to_dtype<float32_t>()}); auto bias = Tensor::from_values({1.f}, {1, 1}, device_t::CPU); Conv2d conv(Conv2dOptions(1, 1, 2).with_bias(true)); conv->set_weights(kernels); conv->set_bias(bias); auto res = conv(im); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 0, 0}), Catch::WithinRel(43.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 0, 1}), Catch::WithinRel(62.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 1, 0}), Catch::WithinRel(100.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 1, 1}), Catch::WithinRel(119.f, 0.001f)); mean(res).backward(); REQUIRE_THAT(conv->weights().grad().at<float32_t>({0, 0, 0, 0}), Catch::WithinRel(3.f, 0.001f)); REQUIRE_THAT(conv->weights().grad().at<float32_t>({0, 0, 0, 1}), Catch::WithinRel(4.f, 0.001f)); REQUIRE_THAT(conv->weights().grad().at<float32_t>({0, 0, 1, 0}), Catch::WithinRel(6.f, 0.001f)); REQUIRE_THAT(conv->weights().grad().at<float32_t>({0, 0, 1, 1}), Catch::WithinRel(7.f, 0.001f)); REQUIRE_THAT(conv->bias().grad().at<float32_t>({0, 0}), Catch::WithinRel(1.f, 0.001f)); } TEST_CASE("Conv2d im3x3 k2x2 ch1x2 bias", "[conv2d]") { std::vector<float32_t> im_data = {1, 2, 3, 4, 5, 6, 7, 8, 9}; std::vector<float32_t> kernel_data = {7, 8, 1, 3, 7, 8, 1, 3}; auto im = Tensor::from_blob( im_data.data(), {.shape = {1, 1, 3, 3}, .device = device_t::CPU, .dtype = to_dtype<float32_t>()}); auto kernels = Tensor::from_blob( kernel_data.data(), {.shape = {2, 1, 2, 2}, .device = device_t::CPU, .dtype = to_dtype<float32_t>()}); auto bias = Tensor::from_values({1.f, 1.f}, {1, 2}, device_t::CPU); Conv2d conv(Conv2dOptions(1, 2, 2).with_bias(true)); conv->set_weights(kernels); conv->set_bias(bias); auto res = conv(im); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 0, 0}), Catch::WithinRel(43.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 0, 1}), Catch::WithinRel(62.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 1, 0}), Catch::WithinRel(100.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 0, 1, 1}), Catch::WithinRel(119.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 1, 0, 0}), Catch::WithinRel(43.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 1, 0, 1}), Catch::WithinRel(62.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 1, 1, 0}), Catch::WithinRel(100.f, 0.001f)); REQUIRE_THAT(res.data().at<float32_t>({0, 1, 1, 1}), Catch::WithinRel(119.f, 0.001f)); }