Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
StaticSymmetry.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2013 Christian Seiler <christian@iwakd.de>
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_TENSORSYMMETRY_STATICSYMMETRY_H
12#define EIGEN_TENSORSYMMETRY_STATICSYMMETRY_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21template <typename indices_, int flags_>
22struct tensor_static_symgroup_element {
23 typedef indices_ indices;
24 constexpr static int flags = flags_;
25};
26
27template <typename Gen, int... indices>
28constexpr std::integer_sequence<int,
29 ((indices == Gen::One) ? Gen::Two : ((indices == Gen::Two) ? Gen::One : indices))...>
30tensor_static_symgroup_swapped_indices(std::integer_sequence<int, indices...>) {
31 return {};
32}
33
34template <typename Gen, int N>
35struct tensor_static_symgroup_element_ctor {
36 typedef tensor_static_symgroup_element<
37 decltype(tensor_static_symgroup_swapped_indices<Gen>(std::make_integer_sequence<int, N>{})), Gen::Flags>
38 type;
39};
40
41template <int N>
42struct tensor_static_symgroup_identity_ctor {
43 typedef tensor_static_symgroup_element<std::make_integer_sequence<int, N>, 0> type;
44};
45
46template <typename iib>
47struct tensor_static_symgroup_multiply_helper {
48 template <int... iia>
49 constexpr static std::integer_sequence<int, get<iia, iib>::value...> helper(std::integer_sequence<int, iia...>) {
50 return {};
51 }
52};
53
54template <typename A, typename B>
55struct tensor_static_symgroup_multiply {
56 private:
57 typedef typename A::indices iia;
58 typedef typename B::indices iib;
59 constexpr static int ffa = A::flags;
60 constexpr static int ffb = B::flags;
61
62 public:
63 static_assert(iia::size() == iib::size(), "Cannot multiply symmetry elements with different number of indices.");
64
65 typedef tensor_static_symgroup_element<decltype(tensor_static_symgroup_multiply_helper<iib>::helper(iia())),
66 ffa ^ ffb>
67 type;
68};
69
70template <typename A, typename B>
71struct tensor_static_symgroup_equality {
72 typedef typename A::indices iia;
73 typedef typename B::indices iib;
74 constexpr static int ffa = A::flags;
75 constexpr static int ffb = B::flags;
76 static_assert(iia::size() == iib::size(), "Cannot compare symmetry elements with different number of indices.");
77
78 constexpr static bool value = std::is_same<iia, iib>::value;
79
80 private:
81 /* this should be zero if they are identical, or else the tensor
82 * will be forced to be pure real, pure imaginary or even pure zero
83 */
84 constexpr static int flags_cmp_ = ffa ^ ffb;
85
86 /* either they are not equal, then we don't care whether the flags
87 * match, or they are equal, and then we have to check
88 */
89 constexpr static bool is_zero = value && flags_cmp_ == NegationFlag;
90 constexpr static bool is_real = value && flags_cmp_ == ConjugationFlag;
91 constexpr static bool is_imag = value && flags_cmp_ == (NegationFlag | ConjugationFlag);
92
93 public:
94 constexpr static int global_flags =
95 (is_real ? GlobalRealFlag : 0) | (is_imag ? GlobalImagFlag : 0) | (is_zero ? GlobalZeroFlag : 0);
96};
97
98template <std::size_t NumIndices, typename... Gen>
99struct tensor_static_symgroup {
100 typedef StaticSGroup<Gen...> type;
101 constexpr static std::size_t size = type::static_size;
102};
103
104template <typename Index, std::size_t N, int... ii, int... jj>
105constexpr static std::array<Index, N> tensor_static_symgroup_index_permute(const std::array<Index, N>& idx,
106 std::integer_sequence<int, ii...>,
107 std::integer_sequence<int, jj...>) {
108 return {{idx[ii]..., idx[sizeof...(ii) + jj]...}};
109}
110
111template <typename Index, int... ii>
112static inline std::vector<Index> tensor_static_symgroup_index_permute(const std::vector<Index>& idx,
113 std::integer_sequence<int, ii...>) {
114 std::vector<Index> result{{idx[ii]...}};
115 std::size_t target_size = idx.size();
116 for (std::size_t i = result.size(); i < target_size; i++) result.push_back(idx[i]);
117 return result;
118}
119
120template <typename T>
121struct tensor_static_symgroup_do_apply;
122
123template <typename first, typename... next>
124struct tensor_static_symgroup_do_apply<internal::type_list<first, next...>> {
125 template <typename Op, typename RV, std::size_t SGNumIndices, typename Index, std::size_t NumIndices,
126 typename... Args>
127 static inline RV run(const std::array<Index, NumIndices>& idx, RV initial, Args&&... args) {
128 static_assert(NumIndices >= SGNumIndices,
129 "Can only apply symmetry group to objects that have at least the required amount of indices.");
130 initial = Op::run(tensor_static_symgroup_index_permute(
131 idx, typename first::indices(), std::make_integer_sequence<int, NumIndices - SGNumIndices>{}),
132 first::flags, initial, std::forward<Args>(args)...);
133 return tensor_static_symgroup_do_apply<internal::type_list<next...>>::template run<Op, RV, SGNumIndices>(
134 idx, initial, args...);
135 }
136
137 template <typename Op, typename RV, std::size_t SGNumIndices, typename Index, typename... Args>
138 static inline RV run(const std::vector<Index>& idx, RV initial, Args&&... args) {
139 eigen_assert(idx.size() >= SGNumIndices &&
140 "Can only apply symmetry group to objects that have at least the required amount of indices.");
141 initial = Op::run(tensor_static_symgroup_index_permute(idx, typename first::indices()), first::flags, initial,
142 std::forward<Args>(args)...);
143 return tensor_static_symgroup_do_apply<internal::type_list<next...>>::template run<Op, RV, SGNumIndices>(
144 idx, initial, args...);
145 }
146};
147
148template <>
149struct tensor_static_symgroup_do_apply<internal::type_list<>> {
150 template <typename Op, typename RV, std::size_t SGNumIndices, typename Index, std::size_t NumIndices,
151 typename... Args>
152 static inline RV run(const std::array<Index, NumIndices>&, RV initial, Args&&...) {
153 // do nothing
154 return initial;
155 }
156
157 template <typename Op, typename RV, std::size_t SGNumIndices, typename Index, typename... Args>
158 static inline RV run(const std::vector<Index>&, RV initial, Args&&...) {
159 // do nothing
160 return initial;
161 }
162};
163
164} // end namespace internal
165
166template <typename... Gen>
167class StaticSGroup {
168 constexpr static std::size_t NumIndices = internal::tensor_symmetry_num_indices<Gen...>::value;
169 typedef internal::group_theory::enumerate_group_elements<
170 internal::tensor_static_symgroup_multiply, internal::tensor_static_symgroup_equality,
171 typename internal::tensor_static_symgroup_identity_ctor<NumIndices>::type,
172 internal::type_list<typename internal::tensor_static_symgroup_element_ctor<Gen, NumIndices>::type...>>
173 group_elements;
174 typedef typename group_elements::type ge;
175
176 public:
177 constexpr StaticSGroup() = default;
178 constexpr StaticSGroup(const StaticSGroup<Gen...>&) = default;
179 constexpr StaticSGroup(StaticSGroup<Gen...>&&) = default;
180
181 template <typename Op, typename RV, typename Index, std::size_t N, typename... Args>
182 static inline RV apply(const std::array<Index, N>& idx, RV initial, Args&&... args) {
183 return internal::tensor_static_symgroup_do_apply<ge>::template run<Op, RV, NumIndices>(idx, initial, args...);
184 }
185
186 template <typename Op, typename RV, typename Index, typename... Args>
187 static inline RV apply(const std::vector<Index>& idx, RV initial, Args&&... args) {
188 eigen_assert(idx.size() == NumIndices);
189 return internal::tensor_static_symgroup_do_apply<ge>::template run<Op, RV, NumIndices>(idx, initial, args...);
190 }
191
192 constexpr static std::size_t static_size = ge::count;
193
194 constexpr static std::size_t size() { return ge::count; }
195 constexpr static int globalFlags() { return group_elements::global_flags; }
196
197 template <typename Tensor_, typename... IndexTypes>
198 inline internal::tensor_symmetry_value_setter<Tensor_, StaticSGroup<Gen...>> operator()(
199 Tensor_& tensor, typename Tensor_::Index firstIndex, IndexTypes... otherIndices) const {
200 static_assert(sizeof...(otherIndices) + 1 == Tensor_::NumIndices,
201 "Number of indices used to access a tensor coefficient must be equal to the rank of the tensor.");
202 return operator()(tensor, std::array<typename Tensor_::Index, Tensor_::NumIndices>{{firstIndex, otherIndices...}});
203 }
204
205 template <typename Tensor_>
206 inline internal::tensor_symmetry_value_setter<Tensor_, StaticSGroup<Gen...>> operator()(
207 Tensor_& tensor, std::array<typename Tensor_::Index, Tensor_::NumIndices> const& indices) const {
208 return internal::tensor_symmetry_value_setter<Tensor_, StaticSGroup<Gen...>>(tensor, *this, indices);
209 }
210};
211
212} // end namespace Eigen
213
214#endif // EIGEN_TENSORSYMMETRY_STATICSYMMETRY_H
215
216/*
217 * kate: space-indent on; indent-width 2; mixedindent off; indent-mode cstyle;
218 */
Namespace containing all symbols from the Eigen library.