/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/dispatch/poolNd.hpp
50 строк
2 KB
kolkir
MaxPool2d implementation
09 фев 2025, 16:36
09 фев 2025, 16:36
b44f8a7
Код
Авторство
О чём код?
#pragma once #include <adept/dispatch/dispatcher.hpp> #include <adept/param_array.hpp> #include <adept/tensor.hpp> namespace adept { inline Tensor avg_pool2d_fwd(const Tensor& input, const ParamArray<2>& kernel, const ParamArray<2>& stride, const ParamArray<2>& padding) { static auto& func = Dispatcher::instance().find("avg_pool2d_fwd"); return Dispatcher::instance() .call<Tensor, const Tensor&, const ParamArray<2>&, const ParamArray<2>&, const ParamArray<2>&>(func, input, kernel, stride, padding); } inline Tensor avg_pool2d_bwd(const Tensor& out_grad, const Shape& input_shape, const ParamArray<2>& kernel, const ParamArray<2>& stride, const ParamArray<2>& padding) { static auto& func = Dispatcher::instance().find("avg_pool2d_bwd"); return Dispatcher::instance() .call<Tensor, const Tensor&, const Shape&, const ParamArray<2>&, const ParamArray<2>&, const ParamArray<2>&>(func, out_grad, input_shape, kernel, stride, padding); } inline Tensor max_pool2d_fwd(const Tensor& input, Tensor& indices, const ParamArray<2>& kernel, const ParamArray<2>& stride, const ParamArray<2>& padding, const ParamArray<2>& dilation) { static auto& func = Dispatcher::instance().find("max_pool2d_fwd"); return Dispatcher::instance() .call<Tensor, const Tensor&, Tensor&, const ParamArray<2>&, const ParamArray<2>&, const ParamArray<2>&, const ParamArray<2>&>(func, input, indices, kernel, stride, padding, dilation); } inline Tensor max_pool2d_bwd(const Tensor& out_grad, const Tensor& indices, const Shape& input_shape) { static auto& func = Dispatcher::instance().find("max_pool2d_bwd"); return Dispatcher::instance().call<Tensor, const Tensor&, const Tensor&, const Shape&>( func, out_grad, indices, input_shape); } } // namespace adept