6#include <stratax/core/dtypes/Concepts.hpp>
7#include <stratax/exceptions/Exceptions.hpp>
8#include <stratax/core/Shape.hpp>
9#include <stratax/containers/Tensor.hpp>
10#include <stratax/containers/Matrix.hpp>
11#include <stratax/containers/Vector.hpp>
12#include <stratax/core/Slice.hpp>
13#include <stratax/indexing/Slicing.hpp>
14#include <stratax/algorithms/Conversion.hpp>
15#include <stratax/core/ReductionTraits.hpp>
38 std::vector<std::size_t>& indices)
40 for (std::size_t d = shape.
rank(); d-- > 0;)
44 if (indices[d] < shape[d])
66inline int normalize_axis(
const A& arr,
int axis)
89inline std::vector<std::size_t> result_shape(
const A& arr,
int axis,
bool keepdims)
93 std::vector<std::size_t> result_dimensions;
97 for (std::size_t dimension = 0; dimension < input_shape.
rank(); ++dimension)
99 if (dimension ==
static_cast<std::size_t
>(axis))
101 result_dimensions.push_back(1);
106 result_dimensions.push_back(input_shape[dimension]);
113 for (std::size_t dimension = 0; dimension < input_shape.
rank(); ++dimension)
115 if (dimension !=
static_cast<std::size_t
>(axis))
117 result_dimensions.push_back(input_shape[dimension]);
122 return result_dimensions;
131template<Array A,
typename Func>
132using axis_reduce_value_t =
133 decltype(std::declval<Func>()(
134 std::declval<const stratax::container::Tensor<typename A::value_type>&>()));
163template<Array A,
typename Func>
165axis_reduce(
const A& array,
int axis, Func func,
bool keepdims =
false)
167 using ResultType = detail::axis_reduce_value_t<A, Func>;
169 int Axis = detail::normalize_axis(array, axis);
171 if (Axis < 0 || Axis >=
static_cast<int>(array.rank()))
177 stratax::conversion::to_tensor(array);
181 std::vector<std::size_t> result_dims =
182 detail::result_shape(array, Axis, keepdims);
186 if (result_dims.empty())
188 ResultType scalar_result = func(arr);
198 std::vector<std::size_t> output_index(result.rank(), 0);
199 std::vector<stratax::core::Slice> slices;
203 std::size_t output_position = 0;
204 for (std::size_t dimension = 0; dimension < input_shape.
rank(); ++dimension)
206 if (
static_cast<int>(dimension) == Axis)
209 static_cast<std::ptrdiff_t
>(input_shape[dimension])});
214 const std::size_t index = keepdims
215 ? output_index[dimension]
216 : output_index[output_position++];
219 static_cast<std::ptrdiff_t
>(index),
220 static_cast<std::ptrdiff_t
>(index + 1)
224 auto s = stratax::indexing::slice(arr, slices);
225 ResultType value = func(s);
226 result(output_index) = value;
228 while (detail::advance(result.shape(), output_index));
245auto sum(
const A& arr)
248 reduction_sum_t<typename A::value_type>;
250 return std::accumulate(
266auto prod(
const A& arr)
269 reduction_prod_t<typename A::value_type>;
271 return std::accumulate(
275 std::multiplies<result_type>{});
288auto max(
const A& arr)
295 auto result = std::max_element(
313auto min(
const A& arr)
320 auto result = std::min_element(
338auto argmax(
const A& arr)
345 auto result = std::max_element(
350 return static_cast<stratax::dtype::int64
>(std::distance(arr.begin(), result));
363auto argmin(
const A& arr)
370 auto result = std::min_element(
375 return static_cast<stratax::dtype::int64
>(std::distance(arr.begin(), result));
391double mean(
const A& arr)
398 return static_cast<double>(sum(arr)) /
static_cast<double>(arr.size());
414double var(
const A& arr)
422 double mean_value = 0.0;
425 for (
const auto& value : arr)
428 const double delta =
static_cast<double>(value) - mean_value;
429 mean_value += delta / count;
430 const double delta2 =
static_cast<double>(value) - mean_value;
431 m2 += delta * delta2;
450double std(
const A& arr)
452 auto vars = var(arr);
453 return std::sqrt(vars);
470auto sum(
const A& arr,
int axis)
476 return reduction::sum(s);
492auto sum(
const A& arr,
int axis,
bool keepdims)
498 return reduction::sum(s);
514auto prod(
const A& arr,
int axis)
520 return reduction::prod(s);
536auto prod(
const A& arr,
int axis,
bool keepdims)
542 return reduction::prod(s);
557auto max(
const A& arr,
int axis)
560 arr, axis, [](
const auto& s) {
return reduction::max(s); });
574auto max(
const A& arr,
int axis,
bool keepdims)
579 [](
const auto& s) {
return reduction::max(s); },
593auto min(
const A& arr,
int axis)
596 arr, axis, [](
const auto& s) {
return reduction::min(s); });
610auto min(
const A& arr,
int axis,
bool keepdims)
615 [](
const auto& s) {
return reduction::min(s); },
629auto argmax(
const A& arr,
int axis)
632 arr, axis, [](
const auto& s) {
return reduction::argmax(s); });
646auto argmax(
const A& arr,
int axis,
bool keepdims)
651 [](
const auto& s) {
return reduction::argmax(s); },
665auto argmin(
const A& arr,
int axis)
668 arr, axis, [](
const auto& s) {
return reduction::argmin(s); });
682auto argmin(
const A& arr,
int axis,
bool keepdims)
687 [](
const auto& s) {
return reduction::argmin(s); },
706mean(
const A& arr,
int axis,
bool keepdims)
711 [](
const auto& s) {
return reduction::mean(s); },
729mean(
const A& arr,
int axis)
732 arr, axis, [](
const auto& s) {
return reduction::mean(s); });
750var(
const A& arr,
int axis,
bool keepdims)
755 [](
const auto& s) {
return reduction::var(s); },
773var(
const A& arr,
int axis)
776 arr, axis, [](
const auto& s) {
return reduction::var(s); });
794std(
const A& arr,
int axis,
bool keepdims)
799 [](
const auto& s) {
return reduction::std(s); },
817std(
const A& arr,
int axis)
820 arr, axis, [](
const auto& s) {
return reduction::std(s); });
Arbitrary-rank owning array of numeric values.
const Shape & shape() const noexcept
Returns the logical shape metadata.
Stores the dimensions of a multidimensional array.
size_type rank() const noexcept
Returns the number of dimensions.
Describes a signed, strided half-open index range.