iterative-solver 0.0
ArrayHandler.h
1#ifndef LINEARALGEBRA_SRC_MOLPRO_LINALG_ARRAY_ARRAYHANDLER_H
2#define LINEARALGEBRA_SRC_MOLPRO_LINALG_ARRAY_ARRAYHANDLER_H
3#include <algorithm>
4#include <functional>
5#include <list>
6#include <map>
7#include <memory>
8#include <numeric>
9#include <set>
10#include <stdexcept>
11#include <string>
12#include <vector>
13
14#include <molpro/linalg/array/type_traits.h>
15#include <molpro/linalg/itsolv/subspace/Matrix.h>
16#include <molpro/linalg/scalar_traits.h>
17#include <molpro/linalg/itsolv/wrap_util.h>
18
23
24namespace util {
25
26struct ArrayHandlerError : public std::logic_error {
27 using std::logic_error::logic_error;
28};
29
32template <typename... Args>
34 using OP = std::tuple<Args...>;
35 std::list<std::tuple<Args...>> m_register;
36
41 template <int N, class ArgEqual>
42 void push(const Args &...args, ArgEqual equal) {
43 auto &&new_op = OP{args...};
44 auto &ref = std::get<N>(new_op);
45 auto rend = std::find_if(m_register.rbegin(), m_register.rend(), [&ref, &equal](const auto &el) {
46 auto xx = std::get<N>(el);
47 return equal(ref, xx);
48 });
49 auto end_of_group = m_register.end();
50 if (rend != m_register.rend())
51 end_of_group = rend.base();
52 m_register.insert(end_of_group, std::forward<OP>(new_op));
53 }
54
56 void push(const Args &...args) { m_register.push_back({args...}); }
57
58 bool empty() { return m_register.empty(); }
59 void clear() { m_register.clear(); }
60};
61
73template <typename X, typename Y, typename Z, class EqualX, class EqualY, class EqualZ>
74std::tuple<std::vector<std::tuple<size_t, size_t, size_t>>, std::vector<X>, std::vector<Y>, std::vector<Z>>
75remove_duplicates(const std::list<std::tuple<X, Y, Z>> &reg, EqualX equal_x, EqualY equal_y, EqualZ equal_z) {
76 auto n_op = reg.size();
77 std::vector<std::tuple<size_t, size_t, size_t>> op_register;
78 op_register.reserve(n_op);
79 std::vector<X> xx;
80 std::vector<Y> yy;
81 std::vector<Z> zz;
82 for (const auto &op : reg) {
83 auto x = std::get<0>(op);
84 auto y = std::get<1>(op);
85 auto z = std::get<2>(op);
86 auto it_x = std::find_if(cbegin(xx), cend(xx), [&x, &equal_x](const auto &el) { return equal_x(x, el); });
87 auto it_y = std::find_if(cbegin(yy), cend(yy), [&y, &equal_y](const auto &el) { return equal_y(y, el); });
88 auto it_z = std::find_if(cbegin(zz), cend(zz), [&z, &equal_z](const auto &el) { return equal_z(z, el); });
89 auto ix = distance(cbegin(xx), it_x);
90 auto iy = distance(cbegin(yy), it_y);
91 auto iz = distance(cbegin(zz), it_z);
92 if (it_x == cend(xx))
93 xx.push_back(x);
94 if (it_y == cend(yy))
95 yy.push_back(y);
96 if (it_z == cend(zz))
97 zz.push_back(z);
98 op_register.emplace_back(ix, iy, iz);
99 }
100 return {op_register, xx, yy, zz};
101}
102
104template <typename T = int>
105struct RefEqual {
106 bool operator()(const std::reference_wrapper<T> &l, const std::reference_wrapper<T> &r) {
107 return std::addressof(l.get()) == std::addressof(r.get());
108 }
109};
110} // namespace util
111
162template <class AL, class AR = AL>
164protected:
165 ArrayHandler() : m_counter(std::make_unique<Counter>()){};
166 ArrayHandler(const ArrayHandler &) = default;
167
168 struct Counter {
169 int scal = 0;
170 int dot = 0;
171 int axpy = 0;
172 int copy = 0;
173 int gemm_inner = 0;
174 int gemm_outer = 0;
175 };
176
177 std::unique_ptr<Counter> m_counter;
178
179public:
182 using value_type = decltype(value_type_L{} * value_type_R{});
183 using value_type_abs = decltype(check_abs<value_type>());
184
185 virtual AL copy(const AR &source) = 0;
187 virtual void copy(AL &x, const AR &y) = 0;
188 virtual void scal(value_type alpha, AL &x) = 0;
189 virtual void fill(value_type alpha, AL &x) = 0;
190 virtual void axpy(value_type alpha, const AR &x, AL &y) = 0;
198 virtual value_type dot(const AL &x, const AR &y) = 0;
199
203 virtual void gemm_outer(const Matrix<value_type> alphas, const CVecRef<AR> &xx, const VecRef<AL> &yy) = 0;
204
208 virtual Matrix<value_type> gemm_inner(const CVecRef<AL> &xx, const CVecRef<AR> &yy) = 0;
209
220 virtual std::map<size_t, value_type_abs> select_max_dot(size_t n, const AL &x, const AR &y) = 0;
221
230 virtual std::map<size_t, value_type> select(size_t n, const AL &x, bool max = false, bool ignore_sign = false) = 0;
231
232 const Counter &counter() const { return *m_counter; }
233
234 std::string counter_to_string(std::string L, std::string R) {
235 std::string output = "";
236 if (m_counter->scal > 0)
237 output.append(std::to_string(m_counter->scal) + " scaling operations of the " + L + " vectors, ");
238 if (m_counter->copy > 0)
239 output.append(std::to_string(m_counter->copy) + " " + L + "<-" + R + " copy operations, ");
240 if (m_counter->dot > 0)
241 output.append(std::to_string(m_counter->dot) + " dot product operations between the " + L + " and " + R +
242 " vectors, ");
243 if (m_counter->axpy > 0)
244 output.append(std::to_string(m_counter->axpy) + " axpy (" + L + " = a*" + R + " + " + L + ") operations, ");
245 if (m_counter->gemm_inner > 0)
246 output.append(std::to_string(m_counter->gemm_inner) + " gemm_inner operations between the " + L + " and " + R +
247 " vectors, ");
248 if (m_counter->gemm_outer > 0)
249 output.append(std::to_string(m_counter->gemm_outer) + " gemm_outer operations between the " + L + " and " + R +
250 " vectors, ");
251 return output;
252 };
253
255 m_counter->scal = 0;
256 m_counter->copy = 0;
257 m_counter->dot = 0;
258 m_counter->axpy = 0;
259 m_counter->gemm_inner = 0;
260 m_counter->gemm_outer = 0;
261 }
262
264 virtual ~ArrayHandler() {
265 std::for_each(m_lazy_handles.begin(), m_lazy_handles.end(), [](auto &el) {
266 if (auto handle = el.lock())
267 handle->invalidate();
268 });
269 }
270
271protected:
276 virtual void error(const std::string &message) { throw util::ArrayHandlerError{message}; };
277
279 virtual void fused_axpy(const std::vector<std::tuple<size_t, size_t, size_t>> &reg,
280 const std::vector<value_type> &alphas,
281 const std::vector<std::reference_wrapper<const AR>> &xx,
282 std::vector<std::reference_wrapper<AL>> &yy) {
283 for (const auto &i : reg) {
284 size_t ai, xi, yi;
285 std::tie(ai, xi, yi) = i;
286 axpy(alphas[ai], xx[xi].get(), yy[yi].get());
287 }
288 }
289
291 virtual void fused_dot(const std::vector<std::tuple<size_t, size_t, size_t>> &reg,
292 const std::vector<std::reference_wrapper<const AL>> &xx,
293 const std::vector<std::reference_wrapper<const AR>> &yy,
294 std::vector<std::reference_wrapper<value_type>> &out) {
295 for (const auto &i : reg) {
296 size_t xi, yi, zi;
297 std::tie(xi, yi, zi) = i;
298 out[zi].get() = dot(xx[xi].get(), yy[yi].get());
299 }
300 }
301
307 public:
309 template <typename T>
310 using ref_wrap = std::reference_wrapper<T>;
311
312 protected:
314 std::set<std::string> m_op_types;
319
320 void error(std::string message) { m_handler.error(message); };
321
330 virtual bool register_op_type(const std::string &type) {
331 if (m_op_types.count(type) == 0 && !m_op_types.empty())
332 return false;
333 m_op_types.insert(type);
334 return true;
335 }
336
338 void clear() {
339 m_op_types.clear();
340 m_axpy.clear();
341 m_dot.clear();
342 }
343
344 public:
345 explicit LazyHandle(ArrayHandler<AL, AR> &handler) : m_handler{handler} {}
347
348 virtual void axpy(value_type alpha, const AR &x, AL &y) {
349 if (register_op_type("axpy"))
350 m_axpy.push(alpha, std::cref(x), std::ref(y));
351 else
352 error("Failed to register operation type axpy with the current state of the LazyHandle");
353 }
354 virtual void dot(const AL &x, const AR &y, value_type &out) {
355 if (register_op_type("dotLR"))
356 m_dot.push(std::cref(x), std::cref(y), std::ref(out));
357 else
358 error("Failed to register operation type dot with the current state of the LazyHandle");
359 }
360
362 virtual void eval() {
363 if (m_invalid)
364 return;
365 if (!m_axpy.empty()) {
366 auto reg = util::remove_duplicates<value_type, ref_wrap<const AR>, ref_wrap<AL>, std::equal_to<value_type>,
368 m_handler.fused_axpy(std::get<0>(reg), std::get<1>(reg), std::get<2>(reg), std::get<3>(reg));
369 }
370 if (!m_dot.empty()) {
371 auto reg =
372 util::remove_duplicates<ref_wrap<const AL>, ref_wrap<const AR>, ref_wrap<value_type>,
374 m_dot.m_register, {}, {}, {});
375 m_handler.fused_dot(std::get<0>(reg), std::get<1>(reg), std::get<2>(reg), std::get<3>(reg));
376 }
377 clear();
378 }
379
381 void invalidate() { m_invalid = true; }
384 bool invalid() { return m_invalid; }
385
386 protected:
388 bool m_invalid = false;
389 };
390
393 public:
394 ProxyHandle(std::shared_ptr<LazyHandle> handle) : m_lazy_handle{std::move(handle)} {}
395
396 template <typename... Args>
397 void axpy(Args &&...args) {
398 m_lazy_handle->axpy(std::forward<Args>(args)...);
399 if (m_off)
400 eval();
401 }
402 template <typename... Args>
403 void dot(Args &&...args) {
404 m_lazy_handle->dot(std::forward<Args>(args)...);
405 if (m_off)
406 eval();
407 }
408 void eval() { m_lazy_handle->eval(); }
409 void invalidate() { m_lazy_handle->invalidate(); }
410 bool invalid() { return m_lazy_handle->invalid(); }
411
413 void off() { m_off = true; };
415 void on() { m_off = false; };
417 bool is_off() { return m_off; }
418
419 protected:
420 std::shared_ptr<LazyHandle> m_lazy_handle;
421 bool m_off = false;
422 };
423
424 std::vector<std::weak_ptr<LazyHandle>> m_lazy_handles;
425
427 void save_handle(const std::shared_ptr<LazyHandle> &handle) {
428 auto empty_handle =
429 std::find_if(m_lazy_handles.begin(), m_lazy_handles.end(), [](const auto &el) { return el.expired(); });
430 if (empty_handle == m_lazy_handles.end())
431 m_lazy_handles.push_back(handle);
432 else
433 *empty_handle = handle;
434 }
435
437 auto handle = std::make_shared<typename ArrayHandler<AL, AR>::LazyHandle>(handler);
438 save_handle(handle);
439 return handle;
440 };
441
442public:
445};
446
447} // namespace molpro::linalg::array
448
449#endif // LINEARALGEBRA_SRC_MOLPRO_LINALG_ARRAY_ARRAYHANDLER_H
Registers operations for lazy evaluation. Evaluation is triggered by calling eval() or on destruction...
Definition: ArrayHandler.h:306
virtual void axpy(value_type alpha, const AR &x, AL &y)
Definition: ArrayHandler.h:348
virtual void eval()
Calls handler to evaluate the registered operations.
Definition: ArrayHandler.h:362
void clear()
Clear the registry.
Definition: ArrayHandler.h:338
virtual void dot(const AL &x, const AR &y, value_type &out)
Definition: ArrayHandler.h:354
virtual ~LazyHandle()
Definition: ArrayHandler.h:346
void error(std::string message)
Definition: ArrayHandler.h:320
void invalidate()
Flag the handler as invalid so that no new operations are registered operations eval() does nothing.
Definition: ArrayHandler.h:381
util::OperationRegister< ref_wrap< const AL >, ref_wrap< const AR >, ref_wrap< value_type > > m_dot
register of dot operations
Definition: ArrayHandler.h:318
LazyHandle(ArrayHandler< AL, AR > &handler)
Definition: ArrayHandler.h:345
ArrayHandler< AL, AR >::value_type value_type
Definition: ArrayHandler.h:308
std::reference_wrapper< T > ref_wrap
Definition: ArrayHandler.h:310
ArrayHandler< AL, AR > & m_handler
all operations are still done through the handler
Definition: ArrayHandler.h:387
std::set< std::string > m_op_types
Types of operations currently registered. Types are strings, because derived classes might add new op...
Definition: ArrayHandler.h:314
bool m_invalid
flags if the handler has been destroyed and LazyHandle is now invalid
Definition: ArrayHandler.h:388
util::OperationRegister< value_type, ref_wrap< const AR >, ref_wrap< AL > > m_axpy
register of axpy operations
Definition: ArrayHandler.h:316
bool invalid()
Definition: ArrayHandler.h:384
virtual bool register_op_type(const std::string &type)
Register an operation type.
Definition: ArrayHandler.h:330
A convenience wrapper around a pointer to the LazyHandle.
Definition: ArrayHandler.h:392
bool is_off()
Returns true if lazy evaluation is off.
Definition: ArrayHandler.h:417
ProxyHandle(std::shared_ptr< LazyHandle > handle)
Definition: ArrayHandler.h:394
void eval()
Definition: ArrayHandler.h:408
void axpy(Args &&...args)
Definition: ArrayHandler.h:397
void invalidate()
Definition: ArrayHandler.h:409
bool invalid()
Definition: ArrayHandler.h:410
bool m_off
whether lazy evaluation is on or off
Definition: ArrayHandler.h:421
void dot(Args &&...args)
Definition: ArrayHandler.h:403
void on()
Turn on lazy evaluation.
Definition: ArrayHandler.h:415
void off()
Turn off lazy evaluation. Next operation will evaluate without delay.
Definition: ArrayHandler.h:413
std::shared_ptr< LazyHandle > m_lazy_handle
Definition: ArrayHandler.h:420
Enhances various operations between pairs of arrays and allows dynamic code injection with uniform in...
Definition: ArrayHandler.h:163
virtual value_type dot(const AL &x, const AR &y)=0
The hermitian inner product <x|y>: conjugate-linear in x, linear in y.
std::unique_ptr< Counter > m_counter
Definition: ArrayHandler.h:177
virtual void fused_dot(const std::vector< std::tuple< size_t, size_t, size_t > > &reg, const std::vector< std::reference_wrapper< const AL > > &xx, const std::vector< std::reference_wrapper< const AR > > &yy, std::vector< std::reference_wrapper< value_type > > &out)
Default implementation of fused_dot without any simplification.
Definition: ArrayHandler.h:291
virtual std::map< size_t, value_type > select(size_t n, const AL &x, bool max=false, bool ignore_sign=false)=0
Select n indices with largest (or smallest) actual (or absolute) value.
typename array::mapped_or_value_type_t< AR > value_type_R
Definition: ArrayHandler.h:181
std::vector< std::weak_ptr< LazyHandle > > m_lazy_handles
keeps track of all created lazy handles
Definition: ArrayHandler.h:424
virtual void scal(value_type alpha, AL &x)=0
typename array::mapped_or_value_type_t< AL > value_type_L
Definition: ArrayHandler.h:180
virtual void gemm_outer(const Matrix< value_type > alphas, const CVecRef< AR > &xx, const VecRef< AL > &yy)=0
decltype(value_type_L{} *value_type_R{}) value_type
Definition: ArrayHandler.h:182
virtual AL copy(const AR &source)=0
ArrayHandler(const ArrayHandler &)=default
virtual void copy(AL &x, const AR &y)=0
Copy content of y into x.
virtual ProxyHandle lazy_handle()=0
Returns a lazy handle. Most implementations simply need to call the overload: return lazy_handle(*thi...
virtual void axpy(value_type alpha, const AR &x, AL &y)=0
virtual Matrix< value_type > gemm_inner(const CVecRef< AL > &xx, const CVecRef< AR > &yy)=0
virtual void fill(value_type alpha, AL &x)=0
virtual void fused_axpy(const std::vector< std::tuple< size_t, size_t, size_t > > &reg, const std::vector< value_type > &alphas, const std::vector< std::reference_wrapper< const AR > > &xx, std::vector< std::reference_wrapper< AL > > &yy)
Default implementation of fused_axpy without any simplification.
Definition: ArrayHandler.h:279
virtual std::map< size_t, value_type_abs > select_max_dot(size_t n, const AL &x, const AR &y)=0
Select n indices with largest by absolute value contributions to the dot product.
virtual ~ArrayHandler()
Destroys ArrayHandler instance and invalidates any LazyHandler it created. Invalidated handler will n...
Definition: ArrayHandler.h:264
ArrayHandler()
Definition: ArrayHandler.h:165
std::string counter_to_string(std::string L, std::string R)
Definition: ArrayHandler.h:234
decltype(check_abs< value_type >()) value_type_abs
Definition: ArrayHandler.h:183
const Counter & counter() const
Definition: ArrayHandler.h:232
ProxyHandle lazy_handle(ArrayHandler< AL, AR > &handler)
Definition: ArrayHandler.h:436
virtual void error(const std::string &message)
Throws an error.
Definition: ArrayHandler.h:276
void save_handle(const std::shared_ptr< LazyHandle > &handle)
Save weak ptr to a lazy handle.
Definition: ArrayHandler.h:427
void clear_counter()
Definition: ArrayHandler.h:254
Matrix container that allows simple data access, slicing, copying and resizing without loosing data.
Definition: Matrix.h:32
std::tuple< std::vector< std::tuple< size_t, size_t, size_t > >, std::vector< X >, std::vector< Y >, std::vector< Z > > remove_duplicates(const std::list< std::tuple< X, Y, Z > > &reg, EqualX equal_x, EqualY equal_y, EqualZ equal_z)
Find duplicates references to x and y arrays and store unique elements in a separate vector.
Definition: ArrayHandler.h:75
Definition: ArrayHandler.h:19
DistrArrayConstIterator cbegin(const DistrArray &array)
Definition: DistrArray.h:442
typename mapped_or_value_type< A >::type mapped_or_value_type_t
Definition: type_traits.h:37
DistrArrayConstIterator cend(const DistrArray &array)
Definition: DistrArray.h:443
std::vector< std::reference_wrapper< const A > > CVecRef
Definition: wrap.h:14
std::vector< std::reference_wrapper< A > > VecRef
Definition: wrap.h:11
Definition: ArrayHandler.h:168
int dot
Definition: ArrayHandler.h:170
int gemm_inner
Definition: ArrayHandler.h:173
int copy
Definition: ArrayHandler.h:172
int axpy
Definition: ArrayHandler.h:171
int gemm_outer
Definition: ArrayHandler.h:174
int scal
Definition: ArrayHandler.h:169
std::list< std::tuple< Args... > > m_register
ordered register of operations
Definition: ArrayHandler.h:35
void push(const Args &...args, ArgEqual equal)
Definition: ArrayHandler.h:42
void clear()
Definition: ArrayHandler.h:59
bool empty()
Definition: ArrayHandler.h:58
std::tuple< Args... > OP
Definition: ArrayHandler.h:34
void push(const Args &...args)
Register each operation as it comes with no reordering.
Definition: ArrayHandler.h:56
When called returns true if addresses of two references are the same.
Definition: ArrayHandler.h:105
bool operator()(const std::reference_wrapper< T > &l, const std::reference_wrapper< T > &r)
Definition: ArrayHandler.h:106