11#ifndef EIGEN_TENSOR_TENSOR_DIMENSIONS_H
12#define EIGEN_TENSOR_TENSOR_DIMENSIONS_H
15#include "./InternalHeaderCheck.h"
22template <std::ptrdiff_t n,
typename Dimension>
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...>> {};
28template <
typename T, T first, T... rest>
29struct dget<0, std::integer_sequence<T, first, rest...>> {
30 static constexpr T value = first;
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) {
39 : (NumIndices - n) > (indices) + dget <
RowMajor ? n - 1
41 Dimensions > ::value * fixed_size_tensor_index_linearization_helper<Index, NumIndices, n - 1, RowMajor>::run(
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&) {
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);
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&) {
86template <std::ptrdiff_t... Indices>
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);
93 constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::ptrdiff_t rank()
const {
return count; }
95 static constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::ptrdiff_t TotalSize() {
96 return internal::arg_prod(Indices...);
99 constexpr EIGEN_DEVICE_FUNC Sizes() =
default;
100 template <
typename DenseIndex>
101 explicit constexpr EIGEN_DEVICE_FUNC Sizes(
const array<DenseIndex, count>& ) {
104 template <
typename... DenseIndex>
105 constexpr EIGEN_DEVICE_FUNC Sizes(DenseIndex...) {}
106 explicit EIGEN_DEVICE_FUNC Sizes(std::initializer_list<std::ptrdiff_t> ) {
110 template <
typename T>
111 Sizes& operator=(
const T& ) {
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);
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);
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);
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;
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) {
146 : (NumIndices - n - 1) > (indices) + array_get <
RowMajor
148 : (NumIndices - n - 1) >
149 (dimensions)*tensor_index_linearization_helper<Index, NumIndices, n - 1, RowMajor>::run(
150 indices, dimensions);
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);
174template <
typename DenseIndex,
int NumDims>
175struct DSizes : array<DenseIndex, NumDims> {
176 typedef array<DenseIndex, NumDims> Base;
177 static constexpr int count = NumDims;
179 constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index rank()
const {
return NumDims; }
181 constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE DenseIndex TotalSize()
const {
182 return (NumDims == 0) ? 1 : internal::array_prod(*
static_cast<const Base*
>(
this));
185 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE DSizes() {
186 for (
int i = 0; i < NumDims; ++i) {
190 EIGEN_DEVICE_FUNC
explicit DSizes(
const array<DenseIndex, NumDims>& a) : Base(a) {}
192 EIGEN_DEVICE_FUNC
explicit DSizes(
const DenseIndex i0) {
193 eigen_assert(NumDims == 1);
197 EIGEN_DEVICE_FUNC DSizes(
const DimensionList<DenseIndex, NumDims>& a) {
198 for (
int i = 0; i < NumDims; ++i) {
205 template <
typename OtherIndex,
207 std::is_same<DenseIndex, typename internal::promote_index_type<DenseIndex, OtherIndex>::type>::value,
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]);
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,
220 std::is_same<DenseIndex, typename internal::promote_index_type<DenseIndex, OtherIndex>::type>::value,
void*>)
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];
230 template <std::ptrdiff_t... Indices>
231 EIGEN_DEVICE_FUNC DSizes(
const Sizes<Indices...>& a) {
232 for (
int i = 0; i < NumDims; ++i) {
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...));
245 EIGEN_DEVICE_FUNC DSizes& operator=(
const array<DenseIndex, NumDims>& other) {
246 *
static_cast<Base*
>(
this) = other;
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));
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));
261template <
typename IndexType,
int NumDims>
262std::ostream& operator<<(std::ostream& os,
const DSizes<IndexType, NumDims>& dims) {
264 for (
int i = 0; i < NumDims; ++i) {
265 if (i > 0) os <<
", ";
274template <
typename DenseIndex,
int NumDims>
275struct array_size<const DSizes<DenseIndex, NumDims>> {
276 static constexpr ptrdiff_t value = NumDims;
278template <
typename DenseIndex,
int NumDims>
279struct array_size<DSizes<DenseIndex, NumDims>> {
280 static constexpr ptrdiff_t value = NumDims;
282template <std::ptrdiff_t... Indices>
283struct array_size<const Sizes<Indices...>> {
284 static constexpr std::ptrdiff_t value = Sizes<Indices...>::count;
286template <std::ptrdiff_t... Indices>
287struct array_size<Sizes<Indices...>> {
288 static constexpr std::ptrdiff_t value = Sizes<Indices...>::count;
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;
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");
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; }
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);
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; }
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);
Namespace containing all symbols from the Eigen library.