/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
with_cpu
src/dispatch/convNd.hpp
40 строк
2 KB
kolkir
Conv2d backward pass implementation
02 фев 2025, 21:49
02 фев 2025, 21:49
9fe8c9a
Код
Авторство
О чём код?
#pragma once #include <adept/dispatch/dispatcher.hpp> #include <adept/param_array.hpp> #include <adept/tensor.hpp> namespace adept { inline 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) { static auto& func = Dispatcher::instance().find("conv2d_fwd"); return Dispatcher::instance() .call<Tensor, const Tensor&, const Tensor&, const Tensor&, const ParamArray<2>&, const ParamArray<2>&, const ParamArray<2>&, const ParamArray<2>&>( func, input, weight, bias, kernel, stride, padding, dilation); } inline 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) { static auto& func = Dispatcher::instance().find("conv2d_bwd"); Dispatcher::instance() .call<void, const Tensor&, const Tensor&, const Tensor&, Tensor&, Tensor&, Tensor&, const ParamArray<2>&, const ParamArray<2>&, const ParamArray<2>&, const ParamArray<2>&>( func, out_grad, input, weight, input_grad, weight_grad, bias_grad, kernel, stride, padding, dilation); } } // namespace adept