/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
with_cpu
src/nn/cross_entropy.cpp
32 строки
1 KB
kolkir
Variable constructor refactoring + ReLU+ LeakyReLU
11 фев 2025, 21:59
11 фев 2025, 21:59
2d87fb5
Код
Авторство
О чём код?
#include <adept/autograd/stats.hpp> #include <adept/nn/cross_entropy.hpp> #include "../dispatch/log_softmax.hpp" #include "../dispatch/nll_loss.hpp" namespace adept { Variable log_softmax_last_dim(const Variable& logits) { auto result = log_softmax_last_dim_fwd(logits.data()); Variable var(result, {logits}, "log_softmax_last_dim"); var.set_backward_fn([logits = logits, result = result](const auto& out_grad) mutable { if (logits.requires_grad()) logits.add_grad(log_softmax_last_dim_bwd(result, out_grad)); }); return var; } Variable cross_entropy_with_logits(const Variable& logits, const Variable& target) { auto log_prob = log_softmax_last_dim(logits); auto result = nll_loss_fwd(log_prob.data(), target.data()); Variable nll_var(std::move(result), {log_prob}, "cross_entropy_with_logits"); nll_var.set_backward_fn( [log_prob = log_prob, target = target.data()](const auto& out_grad) mutable { if (log_prob.requires_grad()) { auto num_classes = log_prob.data().properties().shape.dim(1); log_prob.add_grad(nll_loss_bwd(target, out_grad, num_classes)); } }); return mean(nll_var); } } // namespace adept