Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorIndexList.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_INDEX_LIST_H
12#define EIGEN_TENSOR_TENSOR_INDEX_LIST_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19template <Index n>
20struct type2index {
21 static constexpr Index value = n;
22 EIGEN_DEVICE_FUNC constexpr operator Index() const { return n; }
23 EIGEN_DEVICE_FUNC void set(Index val) {
24 EIGEN_ONLY_USED_FOR_DEBUG(val);
25 eigen_assert(val == n);
26 }
27};
28
29// This can be used with IndexPairList to get compile-time constant pairs,
30// such as IndexPairList<type2indexpair<1,2>, type2indexpair<3,4>>().
31template <Index f, Index s>
32struct type2indexpair {
33 static constexpr Index first = f;
34 static constexpr Index second = s;
35
36 constexpr EIGEN_DEVICE_FUNC operator IndexPair<Index>() const { return IndexPair<Index>(f, s); }
37
38 EIGEN_DEVICE_FUNC void set(const IndexPair<Index>& val) {
39 EIGEN_ONLY_USED_FOR_DEBUG(val);
40 eigen_assert(val.first == f);
41 eigen_assert(val.second == s);
42 }
43};
44
45template <Index n>
46struct NumTraits<type2index<n>> {
47 typedef Index Real;
48 enum { IsComplex = 0, RequireInitialization = false, ReadCost = 1, AddCost = 1, MulCost = 1 };
49
50 EIGEN_DEVICE_FUNC static constexpr EIGEN_STRONG_INLINE Real epsilon() { return 0; }
51 EIGEN_DEVICE_FUNC static constexpr EIGEN_STRONG_INLINE Real dummy_precision() { return 0; }
52 EIGEN_DEVICE_FUNC static constexpr EIGEN_STRONG_INLINE Real highest() { return n; }
53 EIGEN_DEVICE_FUNC static constexpr EIGEN_STRONG_INLINE Real lowest() { return n; }
54};
55
56namespace internal {
57template <typename T>
58EIGEN_DEVICE_FUNC void update_value(T& val, Index new_val) {
59 val = internal::convert_index<T>(new_val);
60}
61template <Index n>
62EIGEN_DEVICE_FUNC void update_value(type2index<n>& val, Index new_val) {
63 val.set(new_val);
64}
65
66template <typename T>
67EIGEN_DEVICE_FUNC void update_value(T& val, IndexPair<Index> new_val) {
68 val = new_val;
69}
70template <Index f, Index s>
71EIGEN_DEVICE_FUNC void update_value(type2indexpair<f, s>& val, IndexPair<Index> new_val) {
72 val.set(new_val);
73}
74
75template <typename T>
76struct is_compile_time_constant_impl : std::false_type {};
77
78template <Index f, Index s>
79struct is_compile_time_constant_impl<type2indexpair<f, s>> : std::true_type {};
80
81template <Index idx>
82struct is_compile_time_constant_impl<type2index<idx>> : std::true_type {};
83
84template <typename T>
85using is_compile_time_constant = is_compile_time_constant_impl<std::remove_cv_t<std::remove_reference_t<T>>>;
86
87template <typename... T>
88struct IndexTuple;
89
90template <typename T, typename... O>
91struct IndexTuple<T, O...> {
92 EIGEN_DEVICE_FUNC constexpr IndexTuple() = default;
93 EIGEN_DEVICE_FUNC constexpr IndexTuple(const T& v, const O... o) : head(v), others(o...) {}
94
95 static constexpr int count = 1 + sizeof...(O);
96 T head{};
97 IndexTuple<O...> others{};
98 typedef T Head;
99 typedef IndexTuple<O...> Other;
100};
101
102template <typename T>
103struct IndexTuple<T> {
104 EIGEN_DEVICE_FUNC constexpr IndexTuple() = default;
105 EIGEN_DEVICE_FUNC constexpr IndexTuple(const T& v) : head(v) {}
106
107 constexpr static int count = 1;
108 T head{};
109 typedef T Head;
110};
111
112template <int N, typename... T>
113struct IndexTupleExtractor;
114
115template <int N, typename T, typename... O>
116struct IndexTupleExtractor<N, T, O...> {
117 typedef typename IndexTupleExtractor<N - 1, O...>::ValType ValType;
118
119 EIGEN_DEVICE_FUNC static constexpr ValType& get_val(IndexTuple<T, O...>& val) {
120 return IndexTupleExtractor<N - 1, O...>::get_val(val.others);
121 }
122
123 EIGEN_DEVICE_FUNC static constexpr const ValType& get_val(const IndexTuple<T, O...>& val) {
124 return IndexTupleExtractor<N - 1, O...>::get_val(val.others);
125 }
126 template <typename V>
127 EIGEN_DEVICE_FUNC static void set_val(IndexTuple<T, O...>& val, V& new_val) {
128 IndexTupleExtractor<N - 1, O...>::set_val(val.others, new_val);
129 }
130};
131
132template <typename T, typename... O>
133struct IndexTupleExtractor<0, T, O...> {
134 typedef T ValType;
135
136 EIGEN_DEVICE_FUNC static constexpr ValType& get_val(IndexTuple<T, O...>& val) { return val.head; }
137 EIGEN_DEVICE_FUNC static constexpr const ValType& get_val(const IndexTuple<T, O...>& val) { return val.head; }
138 template <typename V>
139 EIGEN_DEVICE_FUNC static void set_val(IndexTuple<T, O...>& val, V& new_val) {
140 val.head = new_val;
141 }
142};
143
144template <int N, typename T, typename... O>
145EIGEN_DEVICE_FUNC constexpr typename IndexTupleExtractor<N, T, O...>::ValType& array_get(IndexTuple<T, O...>& tuple) {
146 return IndexTupleExtractor<N, T, O...>::get_val(tuple);
147}
148template <int N, typename T, typename... O>
149EIGEN_DEVICE_FUNC constexpr const typename IndexTupleExtractor<N, T, O...>::ValType& array_get(
150 const IndexTuple<T, O...>& tuple) {
151 return IndexTupleExtractor<N, T, O...>::get_val(tuple);
152}
153template <typename T, typename... O>
154struct array_size<IndexTuple<T, O...>> {
155 static constexpr size_t value = IndexTuple<T, O...>::count;
156};
157template <typename T, typename... O>
158struct array_size<const IndexTuple<T, O...>> {
159 static constexpr size_t value = IndexTuple<T, O...>::count;
160};
161
162template <Index Idx, typename ValueT>
163struct tuple_coeff {
164 template <typename... T>
165 EIGEN_DEVICE_FUNC static constexpr ValueT get(const Index i, const IndexTuple<T...>& t) {
166 return i == Idx ? array_get<Idx>(t) : tuple_coeff<Idx - 1, ValueT>::get(i, t);
167 }
168 template <typename... T>
169 EIGEN_DEVICE_FUNC static void set(const Index i, IndexTuple<T...>& t, const ValueT& value) {
170 if (i == Idx) {
171 update_value(array_get<Idx>(t), value);
172 } else {
173 tuple_coeff<Idx - 1, ValueT>::set(i, t, value);
174 }
175 }
176
177 template <typename... T>
178 EIGEN_DEVICE_FUNC static constexpr bool value_known_statically(const Index i, const IndexTuple<T...>& t) {
179 return ((i == Idx) && is_compile_time_constant<typename IndexTupleExtractor<Idx, T...>::ValType>::value) ||
180 tuple_coeff<Idx - 1, ValueT>::value_known_statically(i, t);
181 }
182
183 template <typename... T>
184 EIGEN_DEVICE_FUNC static constexpr bool values_up_to_known_statically(const IndexTuple<T...>& t) {
185 return is_compile_time_constant<typename IndexTupleExtractor<Idx, T...>::ValType>::value &&
186 tuple_coeff<Idx - 1, ValueT>::values_up_to_known_statically(t);
187 }
188
189 template <typename... T>
190 EIGEN_DEVICE_FUNC static constexpr bool values_up_to_statically_known_to_increase(const IndexTuple<T...>& t) {
191 return is_compile_time_constant<typename IndexTupleExtractor<Idx, T...>::ValType>::value &&
192 is_compile_time_constant<typename IndexTupleExtractor<Idx, T...>::ValType>::value &&
193 array_get<Idx>(t) > array_get<Idx - 1>(t) &&
194 tuple_coeff<Idx - 1, ValueT>::values_up_to_statically_known_to_increase(t);
195 }
196};
197
198template <typename ValueT>
199struct tuple_coeff<0, ValueT> {
200 template <typename... T>
201 EIGEN_DEVICE_FUNC static constexpr ValueT get(const Index /*i*/, const IndexTuple<T...>& t) {
202 return array_get<0>(t);
203 }
204 template <typename... T>
205 EIGEN_DEVICE_FUNC static void set(const Index i, IndexTuple<T...>& t, const ValueT value) {
206 EIGEN_ONLY_USED_FOR_DEBUG(i);
207 eigen_assert(i == 0);
208 update_value(array_get<0>(t), value);
209 }
210 template <typename... T>
211 EIGEN_DEVICE_FUNC static constexpr bool value_known_statically(const Index i, const IndexTuple<T...>&) {
212 return is_compile_time_constant<typename IndexTupleExtractor<0, T...>::ValType>::value && (i == 0);
213 }
214
215 template <typename... T>
216 EIGEN_DEVICE_FUNC static constexpr bool values_up_to_known_statically(const IndexTuple<T...>&) {
217 return is_compile_time_constant<typename IndexTupleExtractor<0, T...>::ValType>::value;
218 }
219
220 template <typename... T>
221 EIGEN_DEVICE_FUNC static constexpr bool values_up_to_statically_known_to_increase(const IndexTuple<T...>&) {
222 return true;
223 }
224};
225} // namespace internal
226
242
243template <typename FirstType, typename... OtherTypes>
244struct IndexList : internal::IndexTuple<FirstType, OtherTypes...> {
245 EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC constexpr Index operator[](const Index i) const {
246 return internal::tuple_coeff<internal::array_size<internal::IndexTuple<FirstType, OtherTypes...>>::value - 1,
247 Index>::get(i, *this);
248 }
249 EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC constexpr Index get(const Index i) const {
250 return internal::tuple_coeff<internal::array_size<internal::IndexTuple<FirstType, OtherTypes...>>::value - 1,
251 Index>::get(i, *this);
252 }
253 EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC void set(const Index i, const Index value) {
254 return internal::tuple_coeff<internal::array_size<internal::IndexTuple<FirstType, OtherTypes...>>::value - 1,
255 Index>::set(i, *this, value);
256 }
257
258 EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC constexpr std::size_t size() const { return 1 + sizeof...(OtherTypes); }
259
260 EIGEN_DEVICE_FUNC constexpr IndexList(const internal::IndexTuple<FirstType, OtherTypes...>& other)
261 : internal::IndexTuple<FirstType, OtherTypes...>(other) {}
262 EIGEN_DEVICE_FUNC constexpr IndexList(FirstType& first, OtherTypes... other)
263 : internal::IndexTuple<FirstType, OtherTypes...>(first, other...) {}
264 EIGEN_DEVICE_FUNC constexpr IndexList() = default;
265
266 EIGEN_DEVICE_FUNC constexpr bool value_known_statically(const Index i) const {
267 return internal::tuple_coeff<internal::array_size<internal::IndexTuple<FirstType, OtherTypes...>>::value - 1,
268 Index>::value_known_statically(i, *this);
269 }
270 EIGEN_DEVICE_FUNC constexpr bool all_values_known_statically() const {
271 return internal::tuple_coeff<internal::array_size<internal::IndexTuple<FirstType, OtherTypes...>>::value - 1,
272 Index>::values_up_to_known_statically(*this);
273 }
274
275 EIGEN_DEVICE_FUNC constexpr bool values_statically_known_to_increase() const {
276 return internal::tuple_coeff<internal::array_size<internal::IndexTuple<FirstType, OtherTypes...>>::value - 1,
277 Index>::values_up_to_statically_known_to_increase(*this);
278 }
279};
280
281template <typename FirstType, typename... OtherTypes>
282std::ostream& operator<<(std::ostream& os, const IndexList<FirstType, OtherTypes...>& dims) {
283 os << "[";
284 for (size_t i = 0; i < 1 + sizeof...(OtherTypes); ++i) {
285 if (i > 0) os << ", ";
286 os << dims[i];
287 }
288 os << "]";
289 return os;
290}
291
292template <typename FirstType, typename... OtherTypes>
293constexpr IndexList<FirstType, OtherTypes...> make_index_list(FirstType val1, OtherTypes... other_vals) {
294 return IndexList<FirstType, OtherTypes...>(val1, other_vals...);
295}
296
297template <typename FirstType, typename... OtherTypes>
298struct IndexPairList : internal::IndexTuple<FirstType, OtherTypes...> {
299 EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC constexpr IndexPair<Index> operator[](const Index i) const {
300 return internal::tuple_coeff<internal::array_size<internal::IndexTuple<FirstType, OtherTypes...>>::value - 1,
301 IndexPair<Index>>::get(i, *this);
302 }
303 EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC void set(const Index i, const IndexPair<Index> value) {
304 return internal::tuple_coeff<internal::array_size<internal::IndexTuple<FirstType, OtherTypes...>>::value - 1,
305 IndexPair<Index>>::set(i, *this, value);
306 }
307
308 EIGEN_DEVICE_FUNC constexpr IndexPairList(const internal::IndexTuple<FirstType, OtherTypes...>& other)
309 : internal::IndexTuple<FirstType, OtherTypes...>(other) {}
310 EIGEN_DEVICE_FUNC constexpr IndexPairList() = default;
311
312 EIGEN_DEVICE_FUNC constexpr bool value_known_statically(const Index i) const {
313 return internal::tuple_coeff<internal::array_size<internal::IndexTuple<FirstType, OtherTypes...>>::value - 1,
314 Index>::value_known_statically(i, *this);
315 }
316};
317
318namespace internal {
319
320template <typename FirstType, typename... OtherTypes>
321EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index array_prod(const IndexList<FirstType, OtherTypes...>& sizes) {
322 Index result = 1;
323 EIGEN_UNROLL_LOOP
324 for (size_t i = 0; i < array_size<IndexList<FirstType, OtherTypes...>>::value; ++i) {
325 result *= sizes[i];
326 }
327 return result;
328}
329
330template <typename FirstType, typename... OtherTypes>
331struct array_size<IndexList<FirstType, OtherTypes...>> {
332 static constexpr size_t value = array_size<IndexTuple<FirstType, OtherTypes...>>::value;
333};
334template <typename FirstType, typename... OtherTypes>
335struct array_size<const IndexList<FirstType, OtherTypes...>> {
336 static constexpr size_t value = array_size<IndexTuple<FirstType, OtherTypes...>>::value;
337};
338
339template <typename FirstType, typename... OtherTypes>
340struct array_size<IndexPairList<FirstType, OtherTypes...>> {
341 static constexpr size_t value = 1 + sizeof...(OtherTypes);
342};
343template <typename FirstType, typename... OtherTypes>
344struct array_size<const IndexPairList<FirstType, OtherTypes...>> {
345 static constexpr size_t value = 1 + sizeof...(OtherTypes);
346};
347
348template <Index N, typename FirstType, typename... OtherTypes>
349EIGEN_DEVICE_FUNC constexpr Index array_get(IndexList<FirstType, OtherTypes...>& a) {
350 return IndexTupleExtractor<N, FirstType, OtherTypes...>::get_val(a);
351}
352template <Index N, typename FirstType, typename... OtherTypes>
353EIGEN_DEVICE_FUNC constexpr Index array_get(const IndexList<FirstType, OtherTypes...>& a) {
354 return IndexTupleExtractor<N, FirstType, OtherTypes...>::get_val(a);
355}
356
357template <typename T>
358struct index_known_statically_impl {
359 EIGEN_DEVICE_FUNC static constexpr bool run(const Index) { return false; }
360};
361
362template <typename FirstType, typename... OtherTypes>
363struct index_known_statically_impl<IndexList<FirstType, OtherTypes...>> {
364 EIGEN_DEVICE_FUNC static constexpr bool run(const Index i) {
365 return IndexList<FirstType, OtherTypes...>().value_known_statically(i);
366 }
367};
368
369template <typename T>
370struct all_indices_known_statically_impl {
371 static constexpr bool run() { return false; }
372};
373
374template <typename FirstType, typename... OtherTypes>
375struct all_indices_known_statically_impl<IndexList<FirstType, OtherTypes...>> {
376 EIGEN_DEVICE_FUNC static constexpr bool run() {
377 return IndexList<FirstType, OtherTypes...>().all_values_known_statically();
378 }
379};
380
381template <typename T>
382struct indices_statically_known_to_increase_impl {
383 EIGEN_DEVICE_FUNC static constexpr bool run() { return false; }
384};
385
386template <typename FirstType, typename... OtherTypes>
387struct indices_statically_known_to_increase_impl<IndexList<FirstType, OtherTypes...>> {
388 EIGEN_DEVICE_FUNC static constexpr bool run() {
389 return Eigen::IndexList<FirstType, OtherTypes...>().values_statically_known_to_increase();
390 }
391};
392
393template <typename Tx>
394struct index_statically_eq_impl {
395 EIGEN_DEVICE_FUNC static constexpr bool run(Index, Index) { return false; }
396};
397
398template <typename FirstType, typename... OtherTypes>
399struct index_statically_eq_impl<IndexList<FirstType, OtherTypes...>> {
400 EIGEN_DEVICE_FUNC static constexpr bool run(const Index i, const Index value) {
401 return IndexList<FirstType, OtherTypes...>().value_known_statically(i) &&
402 (IndexList<FirstType, OtherTypes...>().get(i) == value);
403 }
404};
405
406template <typename T>
407struct index_statically_ne_impl {
408 EIGEN_DEVICE_FUNC static constexpr bool run(Index, Index) { return false; }
409};
410
411template <typename FirstType, typename... OtherTypes>
412struct index_statically_ne_impl<IndexList<FirstType, OtherTypes...>> {
413 EIGEN_DEVICE_FUNC static constexpr bool run(const Index i, const Index value) {
414 return IndexList<FirstType, OtherTypes...>().value_known_statically(i) &&
415 (IndexList<FirstType, OtherTypes...>().get(i) != value);
416 }
417};
418
419template <typename T>
420struct index_statically_gt_impl {
421 EIGEN_DEVICE_FUNC static constexpr bool run(Index, Index) { return false; }
422};
423
424template <typename FirstType, typename... OtherTypes>
425struct index_statically_gt_impl<IndexList<FirstType, OtherTypes...>> {
426 EIGEN_DEVICE_FUNC static constexpr bool run(const Index i, const Index value) {
427 return IndexList<FirstType, OtherTypes...>().value_known_statically(i) &&
428 (IndexList<FirstType, OtherTypes...>().get(i) > value);
429 }
430};
431
432template <typename T>
433struct index_statically_lt_impl {
434 EIGEN_DEVICE_FUNC static constexpr bool run(Index, Index) { return false; }
435};
436
437template <typename FirstType, typename... OtherTypes>
438struct index_statically_lt_impl<IndexList<FirstType, OtherTypes...>> {
439 EIGEN_DEVICE_FUNC static constexpr bool run(const Index i, const Index value) {
440 return IndexList<FirstType, OtherTypes...>().value_known_statically(i) &&
441 (IndexList<FirstType, OtherTypes...>().get(i) < value);
442 }
443};
444
445template <typename Tx>
446struct index_pair_first_statically_eq_impl {
447 EIGEN_DEVICE_FUNC static constexpr bool run(Index, Index) { return false; }
448};
449
450template <typename FirstType, typename... OtherTypes>
451struct index_pair_first_statically_eq_impl<IndexPairList<FirstType, OtherTypes...>> {
452 EIGEN_DEVICE_FUNC static constexpr bool run(const Index i, const Index value) {
453 return IndexPairList<FirstType, OtherTypes...>().value_known_statically(i) &&
454 (IndexPairList<FirstType, OtherTypes...>().operator[](i).first == value);
455 }
456};
457
458template <typename Tx>
459struct index_pair_second_statically_eq_impl {
460 EIGEN_DEVICE_FUNC static constexpr bool run(Index, Index) { return false; }
461};
462
463template <typename FirstType, typename... OtherTypes>
464struct index_pair_second_statically_eq_impl<IndexPairList<FirstType, OtherTypes...>> {
465 EIGEN_DEVICE_FUNC static constexpr bool run(const Index i, const Index value) {
466 return IndexPairList<FirstType, OtherTypes...>().value_known_statically(i) &&
467 (IndexPairList<FirstType, OtherTypes...>().operator[](i).second == value);
468 }
469};
470
471} // end namespace internal
472} // end namespace Eigen
473
474namespace Eigen {
475namespace internal {
476template <typename T>
477static EIGEN_DEVICE_FUNC constexpr bool index_known_statically(Index i) {
478 return index_known_statically_impl<std::remove_cv_t<T>>::run(i);
479}
480
481template <typename T>
482static EIGEN_DEVICE_FUNC constexpr bool all_indices_known_statically() {
483 return all_indices_known_statically_impl<std::remove_cv_t<T>>::run();
484}
485
486template <typename T>
487static EIGEN_DEVICE_FUNC constexpr bool indices_statically_known_to_increase() {
488 return indices_statically_known_to_increase_impl<std::remove_cv_t<T>>::run();
489}
490
491template <typename T>
492static EIGEN_DEVICE_FUNC constexpr bool index_statically_eq(Index i, Index value) {
493 return index_statically_eq_impl<std::remove_cv_t<T>>::run(i, value);
494}
495
496template <typename T>
497static EIGEN_DEVICE_FUNC constexpr bool index_statically_ne(Index i, Index value) {
498 return index_statically_ne_impl<std::remove_cv_t<T>>::run(i, value);
499}
500
501template <typename T>
502static EIGEN_DEVICE_FUNC constexpr bool index_statically_gt(Index i, Index value) {
503 return index_statically_gt_impl<std::remove_cv_t<T>>::run(i, value);
504}
505
506template <typename T>
507static EIGEN_DEVICE_FUNC constexpr bool index_statically_lt(Index i, Index value) {
508 return index_statically_lt_impl<std::remove_cv_t<T>>::run(i, value);
509}
510
511template <typename T>
512static EIGEN_DEVICE_FUNC constexpr bool index_pair_first_statically_eq(Index i, Index value) {
513 return index_pair_first_statically_eq_impl<std::remove_cv_t<T>>::run(i, value);
514}
515
516template <typename T>
517static EIGEN_DEVICE_FUNC constexpr bool index_pair_second_statically_eq(Index i, Index value) {
518 return index_pair_second_statically_eq_impl<std::remove_cv_t<T>>::run(i, value);
519}
520
521} // end namespace internal
522} // end namespace Eigen
523
524#endif // EIGEN_TENSOR_TENSOR_INDEX_LIST_H
Namespace containing all symbols from the Eigen library.