/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/backends/cpu/conv2d.cpp
256 строк
9 KB
kolkir
Fix training issues
06 мар 2025, 23:54
06 мар 2025, 23:54
efe26e8
Код
Авторство
О чём код?
#include <adept/backends/cpu/conv2d.hpp> #include <adept/irange.hpp> #include <adept/types_dispatch.hpp> #include "im2col.hpp" #include "matmul.hpp" namespace adept::cpu { namespace { std::vector<index_t> conv2d_out_size(const Shape& input_size, const ParamArray<2>& kernel_size, const ParamArray<2>& stride_size, const ParamArray<2>& padding_size, const ParamArray<2>& dilation_size) { std::vector<index_t> sizes; for (size_t i = 0; i < 2; ++i) { sizes.push_back((input_size.dim(i + input_size.rang() - 2) + 2 * padding_size[i] - (dilation_size[i] * (kernel_size[i] - 1)) - 1) / stride_size[i] + 1); } return sizes; } template <typename DataType> struct BatchSlicer { template <typename T> BatchSlicer(T& tensor) { if (tensor.defined()) { if constexpr (std::is_const_v<T>) { data_ptr_base = tensor.template const_data_ptr<std::remove_const_t<DataType>>(); } else { data_ptr_base = tensor.template mutable_data_ptr<std::remove_const_t<DataType>>(); } has_batch = tensor.shape().rang() == 4; if (has_batch) { coodrs.resize(tensor.shape().rang()); indexer = Indexer(&tensor.shape()); } } } DataType* get_batch_ptr(index_t i) { if (data_ptr_base != nullptr) { if (!has_batch) { return data_ptr_base; } coodrs[0] = i; auto offset = indexer.value().idxravel(coodrs); return data_ptr_base + offset; } return nullptr; } std::vector<index_t> coodrs; std::optional<Indexer> indexer; DataType* data_ptr_base{nullptr}; bool has_batch{false}; }; } // namespace Tensor conv2d_fwd(const Tensor& input, const Tensor& weight, const Tensor& bias, const ParamArray<2>& kernel, const ParamArray<2>& stride, const ParamArray<2>& padding, const ParamArray<2>& dilation) { // TODO: consider to add check sizes ParamArray<2> input_size{0, 0}; if (input.shape().rang() == 3) { input_size[0] = input.shape().dim(1); input_size[1] = input.shape().dim(2); } else if (input.shape().rang() == 4) { input_size[0] = input.shape().dim(2); input_size[1] = input.shape().dim(3); } else { THROW_ERROR("Invalid input shape ", input.shape(), " for convolution"); } auto is_batch = input.shape().rang() == 4; // NCWH auto output_size = conv2d_out_size(input.shape(), kernel, stride, padding, dilation); index_t batch_size = 0; if (is_batch) batch_size = input.shape().dim(0); const auto in_channels = weight.shape().dim(1); const auto out_channels = weight.shape().dim(0); const auto m = std::accumulate(kernel->begin(), kernel->end(), index_t{1}, std::multiplies<>()); const auto n = std::accumulate(output_size.begin(), output_size.end(), index_t{1}, std::multiplies<>()); auto columns = Tensor::empty({.shape{in_channels * m, n}, .device = input.device(), .dtype = input.dtype()}); auto rear_output_size = output_size; output_size.insert(output_size.begin(), out_channels); if (is_batch) output_size.insert(output_size.begin(), batch_size); auto output = Tensor::zero({.shape = output_size, .device = input.device(), .dtype = input.dtype()}); DISPATCH_FLOATING_TYPES(input.dtype(), [&]() { BatchSlicer<const scalar_t> in_sclicer(input); BatchSlicer<scalar_t> out_sclicer(output); for (auto b : irange(batch_size)) { auto in_data_ptr = in_sclicer.get_batch_ptr(b); auto out_data_ptr = out_sclicer.get_batch_ptr(b); if (bias.defined()) { // fill output with bias values to use gemm addition and skip one tensor allocation index_t channel_size = rear_output_size[0] * rear_output_size[1]; for (auto c : irange(out_channels)) { auto value = bias.at<scalar_t>({0, c}); std::fill_n(out_data_ptr + c * channel_size, channel_size, value); } } im2col(in_data_ptr, in_channels, input_size, rear_output_size, kernel, stride, padding, dilation, columns.mutable_data_ptr<scalar_t>()); // trick with column major orderring to skip col2im step gemm( /* tr_a=*/CblasNoTrans, /* tr_b=*/CblasNoTrans, /* m=*/columns.shape().dim(1), /* n=*/out_channels, /* k=*/columns.shape().dim(0), /* alpha=*/static_cast<scalar_t>(1), /* A=*/columns.const_data_ptr<scalar_t>(), /* lda=*/columns.shape().dim(1), /* B=*/weight.const_data_ptr<scalar_t>(), /* ldb=*/columns.shape().dim(0), /* beta=*/static_cast<scalar_t>(1), /* C=*/out_data_ptr, /* ldc=*/columns.shape().dim(1), /* order=*/CblasColMajor); } }); return output; } void conv2d_bwd(const Tensor& out_grad, const Tensor& input, const Tensor& weight, Tensor& input_grad, Tensor& weight_grad, Tensor& bias_grad, const ParamArray<2>& kernel, const ParamArray<2>& stride, const ParamArray<2>& padding, const ParamArray<2>& dilation) { // TODO: consider to add check sizes ParamArray<2> input_size{0, 0}; if (input.shape().rang() == 3) { input_size[0] = input.shape().dim(1); input_size[1] = input.shape().dim(2); } else if (input.shape().rang() == 4) { input_size[0] = input.shape().dim(2); input_size[1] = input.shape().dim(3); } else { THROW_ERROR("Invalid input shape ", input.shape(), " for convolution"); } auto is_batch = input.shape().rang() == 4; // NCWH auto output_size = conv2d_out_size(input.shape(), kernel, stride, padding, dilation); index_t batch_size = 0; if (is_batch) batch_size = input.shape().dim(0); const auto in_channels = weight.shape().dim(1); const auto out_channels = weight.shape().dim(0); const auto m = std::accumulate(kernel->begin(), kernel->end(), index_t{1}, std::multiplies<>()); const auto n = std::accumulate(output_size.begin(), output_size.end(), index_t{1}, std::multiplies<>()); auto columns = Tensor::empty({.shape{in_channels * m, n}, .device = input.device(), .dtype = input.dtype()}); auto rear_output_size = output_size; output_size.insert(output_size.begin(), out_channels); if (is_batch) output_size.insert(output_size.begin(), batch_size); DISPATCH_FLOATING_TYPES(input.dtype(), [&]() { BatchSlicer<const scalar_t> in_sclicer(input); BatchSlicer<const scalar_t> out_grad_sclicer(out_grad); BatchSlicer<scalar_t> in_grad_sclicer(input_grad); for (auto b : irange(batch_size)) { auto in_data_ptr = in_sclicer.get_batch_ptr(b); auto out_grad_ptr = out_grad_sclicer.get_batch_ptr(b); auto in_grad_ptr = in_grad_sclicer.get_batch_ptr(b); // trick with column major orderring to skip additional computations // Input gradient // -------------------------------------------------------------------------------- if (input_grad.defined()) { gemm( /*transa=*/CblasNoTrans, /*transb=*/CblasTrans, /* m=*/columns.shape().dim(1), /* n=*/columns.shape().dim(0), /* k=*/out_channels, /* alpha=*/static_cast<scalar_t>(1), /* A=*/out_grad_ptr, /* lda=*/columns.shape().dim(1), /* B=*/weight.const_data_ptr<scalar_t>(), /* ldb=*/columns.shape().dim(0), /* beta=*/static_cast<scalar_t>(0), /* C=*/columns.mutable_data_ptr<scalar_t>(), /* ldc=*/columns.shape().dim(1), /* order=*/CblasColMajor); col2im(columns.const_data_ptr<scalar_t>(), in_channels, input_size, rear_output_size, kernel, stride, padding, dilation, in_grad_ptr); } // Weights gradient // ------------------------------------------------------------------------------ im2col(in_data_ptr, in_channels, input_size, rear_output_size, kernel, stride, padding, dilation, columns.mutable_data_ptr<scalar_t>()); gemm( /*transa=*/CblasTrans, /*transb=*/CblasNoTrans, /* m=*/columns.shape().dim(0), /* n=*/out_channels, /* k=*/columns.shape().dim(1), /* alpha=*/static_cast<scalar_t>(1), /* A=*/columns.const_data_ptr<scalar_t>(), /* lda=*/columns.shape().dim(1), /* B=*/out_grad_ptr, /* ldb=*/columns.shape().dim(1), /* beta=*/static_cast<scalar_t>(1), /* C=*/weight_grad.mutable_data_ptr<scalar_t>(), /* ldc=*/columns.shape().dim(0), /* order=*/CblasColMajor); // Bias gradient // ------------------------------------------------------------------------------ if (bias_grad.defined()) { auto bias_grad_ptr = bias_grad.mutable_data_ptr<scalar_t>(); index_t channel_size = rear_output_size[0] * rear_output_size[1]; for (auto c : irange(out_channels)) { auto grad_start = out_grad_ptr + c * channel_size; auto sum = std::accumulate(grad_start, grad_start + channel_size, scalar_t{0}); bias_grad_ptr[c] += sum; } } } }); } } // namespace adept::cpu