3#include <stratax/core/Concepts.hpp>
4#include <stratax/core/Exceptions.hpp>
5#include <stratax/core/containers/Shape.hpp>
7#include <stratax/core/containers/Tensor.hpp>
8#include <stratax/core/containers/Matrix.hpp>
9#include <stratax/core/containers/Vector.hpp>
10#include <stratax/core/Slice.hpp>
11#include <stratax/core/ops/Slice.hpp>
12#include "Conversions.hpp"
30inline bool advance(
const stratax::core::Shape& shape, std::vector<std::size_t>& indices)
33 for (
int d = shape.
rank() - 1; d >= 0; --d) {
35 if (indices[d] < shape(d)) {
54inline int normalize_axis(
const A& arr,
int axis)
76inline std::vector<std::size_t> result_shape(
const A& arr,
int axis,
bool keepdims)
78 stratax::core::Shape input_shape = arr.shape();
80 std::vector<std::size_t> result_dimensions;
84 for (std::size_t dimension = 0; dimension < input_shape.
rank(); ++dimension)
86 if (dimension ==
static_cast<std::size_t
>(axis))
88 result_dimensions.push_back(1);
93 result_dimensions.push_back(input_shape[dimension]);
100 for (std::size_t dimension = 0; dimension < input_shape.
rank(); ++dimension)
102 if (dimension !=
static_cast<std::size_t
>(axis))
104 result_dimensions.push_back(input_shape[dimension]);
109 return result_dimensions;
118template<Array A,
typename Func>
119using axis_reduce_value_t =
120 decltype(std::declval<Func>()(
121 std::declval<const stratax::container::Tensor<typename A::value_type>&>()));
137template<Array A,
typename Func>
138stratax::container::Tensor<axis_reduce_value_t<A, Func>>
139axis_reduce(
const A& array,
int axis, Func func,
bool keepdims =
false)
141 using ResultType = axis_reduce_value_t<A, Func>;
143 int Axis = normalize_axis(array, axis);
145 if (Axis < 0 || Axis >=
static_cast<int>(array.rank()))
147 throw Exceptions::AxisError(
"axis is out of range.");
150 stratax::container::Tensor<typename A::value_type> arr = to_tensor(array);
152 const stratax::core::Shape input_shape = arr.
shape();
154 std::vector<std::size_t> result_dims = result_shape(array, Axis, keepdims);
158 if (result_dims.empty())
160 ResultType scalar_result = func(arr);
161 return stratax::container::Tensor<ResultType>(stratax::core::Shape{1}, scalar_result);
164 stratax::container::Tensor<ResultType> result(stratax::core::Shape{result_dims});
166 std::vector<std::size_t> output_index(result.rank(), 0);
167 std::vector<stratax::core::Slice> slices;
171 std::size_t output_position = 0;
172 for (std::size_t dimension = 0; dimension < input_shape.
rank(); ++dimension)
174 if (
static_cast<int>(dimension) == Axis)
176 slices.push_back(stratax::core::Slice{
static_cast<std::ptrdiff_t
>(0),
177 static_cast<std::ptrdiff_t
>(input_shape[dimension])});
182 const std::size_t index = keepdims
183 ? output_index[dimension]
184 : output_index[output_position++];
186 slices.push_back(stratax::core::Slice{
187 static_cast<std::ptrdiff_t
>(index),
188 static_cast<std::ptrdiff_t
>(index + 1)
192 auto s = slice(arr, slices);
193 ResultType value = func(s);
194 result(output_index) = value;
196 while (reduction::advance(result.shape(), output_index));
212typename A::value_type sum(
const A& arr)
214 return std::accumulate(
217 typename A::value_type(0)
230typename A::value_type prod(
const A& arr)
232 return std::accumulate(
235 typename A::value_type(1),
236 std::multiplies<typename A::value_type>()
249typename A::value_type max(
const A& arr)
251 auto result = std::max_element(
268typename A::value_type min(
const A& arr)
270 auto result = std::min_element(
287std::size_t argmax(
const A& arr)
289 auto result = std::max_element(
294 return static_cast<std::size_t
>(std::distance(arr.begin(), result));
306std::size_t argmin(
const A& arr)
308 auto result = std::min_element(
313 return static_cast<std::size_t
>(std::distance(arr.begin(), result));
325double mean(
const A& arr)
327 return static_cast<double>(sum(arr)) /
static_cast<double>(arr.size());
339double var(
const A& arr)
342 double mean_value = 0.0;
345 for (
const auto& value : arr)
348 const double delta =
static_cast<double>(value) - mean_value;
349 mean_value += delta / count;
350 const double delta2 =
static_cast<double>(value) - mean_value;
351 m2 += delta * delta2;
354 return count > 0.0 ? m2 / count : 0.0;
366double std(
const A& arr)
368 auto vars = var(arr);
369 return std::sqrt(vars);
385stratax::container::Tensor<typename A::value_type> sum(
const A& arr,
int axis)
387 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::sum(s); });
401stratax::container::Tensor<typename A::value_type> sum(
const A& arr,
int axis,
bool keepdims)
403 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::sum(s); }, keepdims);
416stratax::container::Tensor<typename A::value_type> prod(
const A& arr,
int axis)
418 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::prod(s); });
432stratax::container::Tensor<typename A::value_type> prod(
const A& arr,
int axis,
bool keepdims)
434 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::prod(s); }, keepdims);
447stratax::container::Tensor<typename A::value_type> max(
const A& arr,
int axis)
449 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::max(s); });
463stratax::container::Tensor<typename A::value_type> max(
const A& arr,
int axis,
bool keepdims)
465 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::max(s); }, keepdims);
478stratax::container::Tensor<typename A::value_type> min(
const A& arr,
int axis)
480 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::min(s); });
494stratax::container::Tensor<typename A::value_type> min(
const A& arr,
int axis,
bool keepdims)
496 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::min(s); }, keepdims);
509stratax::container::Tensor<std::size_t> argmax(
const A& arr,
int axis)
511 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::argmax(s); });
525stratax::container::Tensor<std::size_t> argmax(
const A& arr,
int axis,
bool keepdims)
527 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::argmax(s); }, keepdims);
540stratax::container::Tensor<std::size_t> argmin(
const A& arr,
int axis)
542 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::argmin(s); });
556stratax::container::Tensor<std::size_t> argmin(
const A& arr,
int axis,
bool keepdims)
558 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::argmin(s); }, keepdims);
572stratax::container::Tensor<double> mean(
const A& arr,
int axis,
bool keepdims)
574 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::mean(s); }, keepdims);
587stratax::container::Tensor<double> mean(
const A& arr,
int axis)
589 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::mean(s); });
603stratax::container::Tensor<double> var(
const A& arr,
int axis,
bool keepdims)
605 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::var(s); }, keepdims);
618stratax::container::Tensor<double> var(
const A& arr,
int axis)
620 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::var(s); });
634stratax::container::Tensor<double> std(
const A& arr,
int axis,
bool keepdims)
636 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::std(s); }, keepdims);
649stratax::container::Tensor<double> std(
const A& arr,
int axis)
651 return axis_reduce(arr, axis, [](
const auto& s) {
return reduction::std(s); });
Shared runtime validation helpers.
const core::Shape & shape() const noexcept
Returns the tensor shape.
std::size_t rank() const
Returns the number of stored dimensions.