8#include <stratax/containers/Matrix.hpp>
9#include <stratax/containers/Tensor.hpp>
10#include <stratax/containers/Vector.hpp>
11#include <stratax/exceptions/Exceptions.hpp>
12#include <stratax/core/Shape.hpp>
13#include <stratax/core/Slice.hpp>
14#include <stratax/core/ArrayView.hpp>
25namespace stratax::indexing {
28using size_type = std::size_t;
30using difference_type = std::ptrdiff_t;
37 difference_type start;
42template<
typename Source>
43using view_element_t = std::conditional_t<
44 std::is_const_v<std::remove_reference_t<Source>>,
45 const typename std::remove_cvref_t<Source>::value_type,
46 typename std::remove_cvref_t<Source>::value_type>;
52 if (extent >
static_cast<size_type
>(std::numeric_limits<difference_type>::max()))
57 const difference_type n =
static_cast<difference_type
>(extent);
58 difference_type start = slice.start();
59 difference_type stop = slice.stop();
60 const difference_type step = slice.step();
73 start = std::clamp(start, difference_type{0}, n);
74 stop = std::clamp(stop, difference_type{0}, n);
78 return ResolvedSlice{start, step, 0};
81 const difference_type distance = stop - start;
82 const size_type count =
static_cast<size_type
>((distance + step - 1) / step);
83 return ResolvedSlice{start, step, count};
90 if (stop < 0 && stop != -1)
95 start = std::clamp(start, difference_type{-1}, n - 1);
96 stop = std::clamp(stop, difference_type{-1}, n - 1);
100 return ResolvedSlice{start, step, 0};
103 const difference_type stride = -step;
104 const difference_type distance = start - stop;
105 const size_type count =
static_cast<size_type
>((distance + stride - 1) / stride);
106 return ResolvedSlice{start, step, count};
113template<
typename Vector>
114auto make_vector_slice_view(
118 const auto resolved =
119 detail::normalize_slice(
123 if (resolved.step < 0)
126 "Negative-step views are not supported.");
130 static_cast<size_type
>(resolved.start);
137 static_cast<size_type
>(resolved.step)
146template<
typename Matrix>
147auto make_matrix_slice_view(
152 const auto resolved_rows =
153 detail::normalize_slice(
157 const auto resolved_cols =
158 detail::normalize_slice(
162 if (resolved_rows.step < 0 || resolved_cols.step < 0)
165 "Negative-step views are not supported.");
168 const size_type offset =
169 static_cast<size_type
>(resolved_rows.start) * mat.strides()[0]
170 +
static_cast<size_type
>(resolved_cols.start) * mat.strides()[1];
178 mat.
strides()[0] *
static_cast<size_type
>(resolved_rows.step),
179 mat.strides()[1] *
static_cast<size_type
>(resolved_cols.step)
195 return detail::make_vector_slice_view(vec, slice);
203 return detail::make_vector_slice_view(vec, slice);
212 return detail::make_matrix_slice_view(mat, rows, cols);
221 return detail::make_matrix_slice_view(mat, rows, cols);
226template<
typename Tensor,
typename... Slices>
229 std::remove_cvref_t<Slices>,
233auto make_tensor_slice_view(
241 if (ranges.size() != tensor.rank())
244 "The number of slices must match the tensor rank.");
252 std::array<size_type,
sizeof...(Slices)> out_dims{};
254 for (size_type dim = 0; dim < ranges.size(); ++dim)
256 resolved[dim] = detail::normalize_slice(
258 tensor.shape()[dim]);
260 out_dims[dim] = resolved[dim].size;
264 std::vector<size_type>(
268 const auto& tensor_strides = tensor.strides();
270 size_type offset = 0;
272 std::vector<size_type> view_stride_values;
273 view_stride_values.reserve(resolved.size());
275 for (size_type dim = 0; dim < resolved.size(); ++dim)
277 if (resolved[dim].step < 0)
280 "Negative-step views are not supported.");
284 static_cast<size_type
>(resolved[dim].start);
286 if (tensor_strides[dim] != 0 &&
287 start > std::numeric_limits<size_type>::max() / tensor_strides[dim])
292 const auto offset_term = start * tensor_strides[dim];
294 if (offset > std::numeric_limits<size_type>::max() - offset_term)
299 offset += offset_term;
301 const auto step =
static_cast<size_type
>(resolved[dim].step);
304 tensor_strides[dim] > std::numeric_limits<size_type>::max() / step)
309 view_stride_values.push_back(tensor_strides[dim] * step);
316 tensor.data() + offset,
321template<
typename Tensor>
322auto make_tensor_slice_view(
324 const std::vector<stratax::core::Slice>& slices)
326 if (slices.size() != tensor.rank())
329 "The number of slices must match the tensor rank.");
332 std::vector<detail::ResolvedSlice> resolved(
335 std::vector<size_type> out_dims(
338 for (size_type dim = 0; dim < slices.size(); ++dim)
340 resolved[dim] = detail::normalize_slice(
342 tensor.shape()[dim]);
344 out_dims[dim] = resolved[dim].size;
347 const auto view_shape =
350 const auto& tensor_strides = tensor.strides();
352 size_type offset = 0;
354 std::vector<size_type> view_stride_values;
355 view_stride_values.reserve(resolved.size());
357 for (size_type dim = 0; dim < resolved.size(); ++dim)
359 if (resolved[dim].step < 0)
362 "Negative-step views are not supported.");
366 static_cast<size_type
>(resolved[dim].start);
368 if (tensor_strides[dim] != 0 &&
369 start > std::numeric_limits<size_type>::max() / tensor_strides[dim])
374 const auto offset_term = start * tensor_strides[dim];
376 if (offset > std::numeric_limits<size_type>::max() - offset_term)
381 offset += offset_term;
383 const auto step =
static_cast<size_type
>(resolved[dim].step);
386 tensor_strides[dim] > std::numeric_limits<size_type>::max() / step)
391 view_stride_values.push_back(tensor_strides[dim] * step);
398 tensor.data() + offset,
405template<
typename T,
typename... Slices>
408 std::remove_cvref_t<Slices>,
414 return detail::make_tensor_slice_view(tensor, slices...);
417template<
typename T,
typename... Slices>
420 std::remove_cvref_t<Slices>,
426 return detail::make_tensor_slice_view(tensor, slices...);
432 const std::vector<stratax::core::Slice>& slices)
434 return detail::make_tensor_slice_view(tensor, slices);
440 const std::vector<stratax::core::Slice>& slices)
442 return detail::make_tensor_slice_view(tensor, slices);
449 const std::vector<stratax::core::Slice>& slices)
451 if (slices.size() != view.rank())
454 "The number of slices must match the view rank.");
457 std::vector<detail::ResolvedSlice> resolved(slices.size());
458 std::vector<size_type> out_dims(slices.size());
459 std::vector<size_type> out_strides(slices.size());
461 size_type start_offset = 0;
463 for (size_type dim = 0; dim < slices.size(); ++dim)
465 resolved[dim] = detail::normalize_slice(
469 out_dims[dim] = resolved[dim].size;
471 if (resolved[dim].step < 0)
474 "Negative-step views are not supported.");
477 const size_type start =
478 static_cast<size_type
>(resolved[dim].start);
480 const size_type step =
481 static_cast<size_type
>(resolved[dim].step);
483 if (view.strides()[dim] != 0 &&
485 std::numeric_limits<size_type>::max() /
489 "ArrayView slice offset overflow.");
492 const size_type start_term =
493 start * view.strides()[dim];
496 std::numeric_limits<size_type>::max() -
500 "ArrayView slice offset overflow.");
503 start_offset += start_term;
506 view.strides()[dim] >
507 std::numeric_limits<size_type>::max() /
511 "ArrayView slice stride overflow.");
515 view.strides()[dim] * step;
521 auto* data = view.data();
523 if (out_shape.elements() != 0)
525 data += start_offset;
Two-dimensional owning array of numeric values.
Arbitrary-rank owning array of numeric values.
One-dimensional owning array of numeric values.
Stores the dimensions of a multidimensional array.
Shape strides() const
Computes canonical row-major strides for this shape.
Describes a signed, strided half-open index range.