/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
with_cpu
src/dispatch/batchnorm.hpp
37 строк
2 KB
kolkir
Batchnorm2d backward implementation
23 фев 2025, 01:20
23 фев 2025, 01:20
d9615da
Код
Авторство
О чём код?
#pragma once #include <adept/dispatch/dispatcher.hpp> #include <adept/tensor.hpp> namespace adept { inline std::tuple<Tensor, Tensor, Tensor> batchnorm2d_fwd(const Tensor& input, const Tensor& weight, const Tensor& bias, Tensor& running_mean, Tensor& running_var, float32_t eps, float32_t momentum, bool is_train) { static auto& func = Dispatcher::instance().find("batchnorm2d_fwd"); return Dispatcher::instance() .call<std::tuple<Tensor, Tensor, Tensor>, const Tensor&, const Tensor&, const Tensor&, Tensor&, Tensor&, float32_t, float32_t, bool>(func, input, weight, bias, running_mean, running_var, eps, momentum, is_train); } inline std::tuple<Tensor, Tensor> batchnorm2d_bwd(Tensor& in_grad, const Tensor& out_grad, const Tensor& input, const Tensor& weight, const Tensor& save_mean, const Tensor& save_var, float32_t eps) { static auto& func = Dispatcher::instance().find("batchnorm2d_bwd"); return Dispatcher::instance() .call<std::tuple<Tensor, Tensor>, Tensor&, const Tensor&, const Tensor&, const Tensor&, const Tensor&, const Tensor&, float32_t>(func, in_grad, out_grad, input, weight, save_mean, save_var, eps); } } // namespace adept