/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
with_cpu
src/backends/cpu/log_softmax.cpp
120 строк
4 KB
kolkir
Refactor simd algos to use map reduce
16 мар 2025, 23:09
16 мар 2025, 23:09
5f536e1
Код
Авторство
О чём код?
#undef HWY_TARGET_INCLUDE #define HWY_TARGET_INCLUDE "../src/backends/cpu/log_softmax.cpp" #include <hwy/foreach_target.h> #include <hwy/highway.h> #include "arithmetics-inl.hpp" #if HWY_ONCE #include <adept/backends/cpu/log_softmax.hpp> #include <adept/print.hpp> #include <adept/threading.hpp> #include <adept/types_dispatch.hpp> #include "arithmetics.hpp" #include <cmath> namespace adept::cpu { /* def log_softmax(x): c = x.max() logsumexp = np.log(np.exp(x - c).sum()) return x - c - logsumexp */ Tensor log_softmax_last_dim_fwd(const Tensor& input) { auto in_shape = input.properties().shape; auto is_shape_aligned = input.is_shape_aligned(); index_t outer_size = 1; index_t dim_size = in_shape.dims().back(); for (index_t i = 0; i < in_shape.rang() - 1; ++i) { outer_size *= in_shape.dim(i); } auto output = input.clone(); DISPATCH_FLOATING_TYPES(input.properties().dtype, [&]() { const scalar_t* input_data_base = input.const_data_ptr<scalar_t>(); scalar_t* output_data_base = output.mutable_data_ptr<scalar_t>(); parallel_for<scalar_t>(0, outer_size, [&](auto begin, auto end) { // rows for (index_t row_i = begin; row_i < end; ++row_i) { const auto* input_data = input_data_base + row_i * dim_size; auto* output_data = output_data_base + row_i * dim_size; // calculate max [scalar] scalar_t max_input = 0; detail::max::apply<scalar_t>(input_data, max_input, dim_size, is_shape_aligned); // calculate tmp[dim_size] = input - max detail::sub::apply<scalar_t>(output_data, max_input, dim_size, is_shape_aligned); // calculate tmp[dim_size] = exp(tmp) detail::exp::apply<scalar_t>(output_data, dim_size, is_shape_aligned); // calculate tmpsum[scalar] = sum(tmp) scalar_t tmpsum = 0; detail::sum::apply<scalar_t>(output_data, tmpsum, dim_size, is_shape_aligned); // calculate tmpsum[scalar] = log(tmpsum) tmpsum = std::log(tmpsum); // calculate out[dimsize] = input[dimsize] - max_input[scalar] - tmpsum[scalar] hwy::CopyBytes(input_data, output_data, dim_size * sizeof(scalar_t)); detail::sub::apply<scalar_t>(output_data, max_input, dim_size, is_shape_aligned); detail::sub::apply<scalar_t>(output_data, tmpsum, dim_size, is_shape_aligned); } }); }); return output; } /* def bwd_log_softmax(res, out_grad): sum = out_grad.sum() return out_grad - np.exp(res) * sum */ Tensor log_softmax_last_dim_bwd(const Tensor& result, const Tensor& out_grad) { auto grad_shape = out_grad.properties().shape; auto is_shape_aligned = out_grad.is_shape_aligned(); index_t outer_size = 1; index_t dim_size = grad_shape.dims().back(); for (index_t i = 0; i < grad_shape.rang() - 1; ++i) { outer_size *= grad_shape.dim(i); } auto input_grad = out_grad.clone(); DISPATCH_FLOATING_TYPES(out_grad.properties().dtype, [&]() { scalar_t* input_grad_data_base = input_grad.mutable_data_ptr<scalar_t>(); const scalar_t* result_data_base = result.const_data_ptr<scalar_t>(); const scalar_t* out_grad_data_base = out_grad.const_data_ptr<scalar_t>(); parallel_for<scalar_t>(0, outer_size, [&](auto begin, auto end) { // row for (index_t row_i = begin; row_i < end; ++row_i) { auto* input_grad_data = input_grad_data_base + row_i * dim_size; const auto* result_data = result_data_base + row_i * dim_size; const auto* out_grad_data = out_grad_data_base + row_i * dim_size; scalar_t sum = 0; detail::sum::apply<scalar_t>(out_grad_data, sum, dim_size, is_shape_aligned); hwy::CopyBytes(result_data, input_grad_data, dim_size * sizeof(scalar_t)); detail::exp::apply<scalar_t>(input_grad_data, dim_size, is_shape_aligned); detail::mul::apply<scalar_t>(input_grad_data, sum, dim_size, is_shape_aligned); detail::neg::apply<scalar_t>(input_grad_data, dim_size, is_shape_aligned); detail::add::apply<scalar_t>(input_grad_data, out_grad_data, dim_size, is_shape_aligned); } }); }); return input_grad; } } // namespace adept::cpu #endif // HWY_ONCE