Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorDimensions.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2014 Benoit Steiner <benoit.steiner.goog@gmail.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_TENSOR_TENSOR_DIMENSIONS_H
12#define EIGEN_TENSOR_TENSOR_DIMENSIONS_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19// Boilerplate code
20namespace internal {
21
22template <std::ptrdiff_t n, typename Dimension>
23struct dget;
24
25template <std::ptrdiff_t n, typename T, T first, T... rest>
26struct dget<n, std::integer_sequence<T, first, rest...>> : dget<n - 1, std::integer_sequence<T, rest...>> {};
27
28template <typename T, T first, T... rest>
29struct dget<0, std::integer_sequence<T, first, rest...>> {
30 static constexpr T value = first;
31};
32
33template <typename Index, std::ptrdiff_t NumIndices, std::ptrdiff_t n, bool RowMajor>
34struct fixed_size_tensor_index_linearization_helper {
35 template <typename Dimensions>
36 constexpr EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Index run(array<Index, NumIndices> const& indices,
37 const Dimensions& dimensions) {
38 return array_get < RowMajor ? n - 1
39 : (NumIndices - n) > (indices) + dget < RowMajor ? n - 1
40 : (NumIndices - n),
41 Dimensions > ::value * fixed_size_tensor_index_linearization_helper<Index, NumIndices, n - 1, RowMajor>::run(
42 indices, dimensions);
43 }
44};
45
46template <typename Index, std::ptrdiff_t NumIndices, bool RowMajor>
47struct fixed_size_tensor_index_linearization_helper<Index, NumIndices, 0, RowMajor> {
48 template <typename Dimensions>
49 constexpr EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Index run(array<Index, NumIndices> const&, const Dimensions&) {
50 return 0;
51 }
52};
53
54template <typename Index, std::ptrdiff_t n>
55struct fixed_size_tensor_index_extraction_helper {
56 template <typename Dimensions>
57 constexpr EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Index run(const Index index, const Dimensions& dimensions) {
58 const Index mult = (index == n - 1) ? 1 : 0;
59 return dget<n - 1, Dimensions>::value * mult +
60 fixed_size_tensor_index_extraction_helper<Index, n - 1>::run(index, dimensions);
61 }
62};
63
64template <typename Index>
65struct fixed_size_tensor_index_extraction_helper<Index, 0> {
66 template <typename Dimensions>
67 constexpr EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Index run(const Index, const Dimensions&) {
68 return 0;
69 }
70};
71
72} // end namespace internal
73
86template <std::ptrdiff_t... Indices>
87struct Sizes {
88 typedef std::integer_sequence<std::ptrdiff_t, Indices...> Base;
89 const Base t = Base();
90 static constexpr std::ptrdiff_t total_size = internal::arg_prod(Indices...);
91 static constexpr ptrdiff_t count = sizeof...(Indices);
92
93 constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::ptrdiff_t rank() const { return count; }
94
95 static constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::ptrdiff_t TotalSize() {
96 return internal::arg_prod(Indices...);
97 }
98
99 constexpr EIGEN_DEVICE_FUNC Sizes() = default;
100 template <typename DenseIndex>
101 explicit constexpr EIGEN_DEVICE_FUNC Sizes(const array<DenseIndex, count>& /*indices*/) {
102 // TODO: Add assertion.
103 }
104 template <typename... DenseIndex>
105 constexpr EIGEN_DEVICE_FUNC Sizes(DenseIndex...) {}
106 explicit EIGEN_DEVICE_FUNC Sizes(std::initializer_list<std::ptrdiff_t> /*l*/) {
107 // TODO: Add assertion.
108 }
109
110 template <typename T>
111 Sizes& operator=(const T& /*other*/) {
112 // add assertion failure if the size of other is different
113 return *this;
114 }
115
116 constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::ptrdiff_t operator[](const std::ptrdiff_t index) const {
117 return internal::fixed_size_tensor_index_extraction_helper<std::ptrdiff_t, count>::run(index, t);
118 }
119
120 template <typename DenseIndex>
121 constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE ptrdiff_t
122 IndexOfColMajor(const array<DenseIndex, count>& indices) const {
123 return internal::fixed_size_tensor_index_linearization_helper<DenseIndex, count, count, false>::run(indices, t);
124 }
125 template <typename DenseIndex>
126 constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE ptrdiff_t
127 IndexOfRowMajor(const array<DenseIndex, count>& indices) const {
128 return internal::fixed_size_tensor_index_linearization_helper<DenseIndex, count, count, true>::run(indices, t);
129 }
130};
131
132namespace internal {
133template <std::ptrdiff_t... Indices>
134constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::ptrdiff_t array_prod(const Sizes<Indices...>&) {
135 return Sizes<Indices...>::total_size;
136}
137} // namespace internal
138
139// Boilerplate
140namespace internal {
141template <typename Index, std::ptrdiff_t NumIndices, std::ptrdiff_t n, bool RowMajor>
142struct tensor_index_linearization_helper {
143 static constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index run(array<Index, NumIndices> const& indices,
144 array<Index, NumIndices> const& dimensions) {
145 return array_get < RowMajor ? n
146 : (NumIndices - n - 1) > (indices) + array_get < RowMajor
147 ? n
148 : (NumIndices - n - 1) >
149 (dimensions)*tensor_index_linearization_helper<Index, NumIndices, n - 1, RowMajor>::run(
150 indices, dimensions);
151 }
152};
153
154template <typename Index, std::ptrdiff_t NumIndices, bool RowMajor>
155struct tensor_index_linearization_helper<Index, NumIndices, 0, RowMajor> {
156 static constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index run(array<Index, NumIndices> const& indices,
157 array<Index, NumIndices> const&) {
158 return array_get < RowMajor ? 0 : NumIndices - 1 > (indices);
159 }
160};
161} // end namespace internal
162
174template <typename DenseIndex, int NumDims>
175struct DSizes : array<DenseIndex, NumDims> {
176 typedef array<DenseIndex, NumDims> Base;
177 static constexpr int count = NumDims;
178
179 constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index rank() const { return NumDims; }
180
181 constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE DenseIndex TotalSize() const {
182 return (NumDims == 0) ? 1 : internal::array_prod(*static_cast<const Base*>(this));
183 }
184
185 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE DSizes() {
186 for (int i = 0; i < NumDims; ++i) {
187 (*this)[i] = 0;
188 }
189 }
190 EIGEN_DEVICE_FUNC explicit DSizes(const array<DenseIndex, NumDims>& a) : Base(a) {}
191
192 EIGEN_DEVICE_FUNC explicit DSizes(const DenseIndex i0) {
193 eigen_assert(NumDims == 1);
194 (*this)[0] = i0;
195 }
196
197 EIGEN_DEVICE_FUNC DSizes(const DimensionList<DenseIndex, NumDims>& a) {
198 for (int i = 0; i < NumDims; ++i) {
199 (*this)[i] = a[i];
200 }
201 }
202
203 // Enable DSizes index type promotion only if we are promoting to the
204 // larger type, e.g. allow to promote dimensions of type int to long.
205 template <typename OtherIndex,
206 std::enable_if_t<
207 std::is_same<DenseIndex, typename internal::promote_index_type<DenseIndex, OtherIndex>::type>::value,
208 int> = 0>
209 EIGEN_DEVICE_FUNC explicit DSizes(const array<OtherIndex, NumDims>& other) {
210 for (int i = 0; i < NumDims; ++i) {
211 (*this)[i] = static_cast<DenseIndex>(other[i]);
212 }
213 }
214
215 template <typename OtherIndex>
216 EIGEN_DEPRECATED_WITH_REASON("Omit the implementation-only second argument.")
217 EIGEN_DEVICE_FUNC explicit DSizes(
218 const array<OtherIndex, NumDims>& other,
219 std::enable_if_t<
220 std::is_same<DenseIndex, typename internal::promote_index_type<DenseIndex, OtherIndex>::type>::value, void*>)
221 : DSizes(other) {}
222
223 template <typename FirstType, typename... OtherTypes>
224 EIGEN_DEVICE_FUNC explicit DSizes(const Eigen::IndexList<FirstType, OtherTypes...>& dimensions) {
225 for (int i = 0; i < dimensions.count; ++i) {
226 (*this)[i] = dimensions[i];
227 }
228 }
229
230 template <std::ptrdiff_t... Indices>
231 EIGEN_DEVICE_FUNC DSizes(const Sizes<Indices...>& a) {
232 for (int i = 0; i < NumDims; ++i) {
233 (*this)[i] = a[i];
234 }
235 }
236
237 template <typename... IndexTypes>
238 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE explicit DSizes(DenseIndex firstDimension, DenseIndex secondDimension,
239 IndexTypes... otherDimensions)
240 : Base({{firstDimension, secondDimension, static_cast<DenseIndex>(otherDimensions)...}}) {
241 EIGEN_STATIC_ASSERT(sizeof...(otherDimensions) + 2 == NumDims, YOU_MADE_A_PROGRAMMING_MISTAKE)
242 eigen_assert(internal::indices_fit<DenseIndex>(otherDimensions...));
243 }
244
245 EIGEN_DEVICE_FUNC DSizes& operator=(const array<DenseIndex, NumDims>& other) {
246 *static_cast<Base*>(this) = other;
247 return *this;
248 }
249
250 // A constexpr would be so much better here
251 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE DenseIndex IndexOfColMajor(const array<DenseIndex, NumDims>& indices) const {
252 return internal::tensor_index_linearization_helper<DenseIndex, NumDims, NumDims - 1, false>::run(
253 indices, *static_cast<const Base*>(this));
254 }
255 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE DenseIndex IndexOfRowMajor(const array<DenseIndex, NumDims>& indices) const {
256 return internal::tensor_index_linearization_helper<DenseIndex, NumDims, NumDims - 1, true>::run(
257 indices, *static_cast<const Base*>(this));
258 }
259};
260
261template <typename IndexType, int NumDims>
262std::ostream& operator<<(std::ostream& os, const DSizes<IndexType, NumDims>& dims) {
263 os << "[";
264 for (int i = 0; i < NumDims; ++i) {
265 if (i > 0) os << ", ";
266 os << dims[i];
267 }
268 os << "]";
269 return os;
270}
271
272namespace internal {
273
274template <typename DenseIndex, int NumDims>
275struct array_size<const DSizes<DenseIndex, NumDims>> {
276 static constexpr ptrdiff_t value = NumDims;
277};
278template <typename DenseIndex, int NumDims>
279struct array_size<DSizes<DenseIndex, NumDims>> {
280 static constexpr ptrdiff_t value = NumDims;
281};
282template <std::ptrdiff_t... Indices>
283struct array_size<const Sizes<Indices...>> {
284 static constexpr std::ptrdiff_t value = Sizes<Indices...>::count;
285};
286template <std::ptrdiff_t... Indices>
287struct array_size<Sizes<Indices...>> {
288 static constexpr std::ptrdiff_t value = Sizes<Indices...>::count;
289};
290template <std::ptrdiff_t n, std::ptrdiff_t... Indices>
291constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::ptrdiff_t array_get(const Sizes<Indices...>&) {
292 return dget<n, typename Sizes<Indices...>::Base>::value;
293}
294template <std::ptrdiff_t n>
295EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::ptrdiff_t array_get(const Sizes<>&) {
296 eigen_assert(false && "should never be called");
297 return -1;
298}
299
300template <typename Dims1, typename Dims2, ptrdiff_t n, ptrdiff_t m>
301struct sizes_match_below_dim {
302 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool run(const Dims1&, const Dims2&) { return false; }
303};
304template <typename Dims1, typename Dims2, ptrdiff_t n>
305struct sizes_match_below_dim<Dims1, Dims2, n, n> {
306 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool run(const Dims1& dims1, const Dims2& dims2) {
307 return numext::equal_strict(array_get<n - 1>(dims1), array_get<n - 1>(dims2)) &&
308 sizes_match_below_dim<Dims1, Dims2, n - 1, n - 1>::run(dims1, dims2);
309 }
310};
311template <typename Dims1, typename Dims2>
312struct sizes_match_below_dim<Dims1, Dims2, 0, 0> {
313 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool run(const Dims1&, const Dims2&) { return true; }
314};
315
316} // end namespace internal
317
318template <typename Dims1, typename Dims2>
319EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool dimensions_match(const Dims1& dims1, const Dims2& dims2) {
320 return internal::sizes_match_below_dim<Dims1, Dims2, internal::array_size<Dims1>::value,
321 internal::array_size<Dims2>::value>::run(dims1, dims2);
322}
323
324} // end namespace Eigen
325
326#endif // EIGEN_TENSOR_TENSOR_DIMENSIONS_H
Namespace containing all symbols from the Eigen library.