Program Listing for File sum.hpp¶
↰ Return to documentation for file (SeQuant/core/expressions/sum.hpp)
#ifndef SEQUANT_EXPRESSIONS_SUM_HPP
#define SEQUANT_EXPRESSIONS_SUM_HPP
#include <SeQuant/core/container.hpp>
#include <SeQuant/core/expressions/constant.hpp>
#include <SeQuant/core/expressions/expr.hpp>
#include <SeQuant/core/expressions/expr_iterator.hpp>
#include <SeQuant/core/expressions/expr_ptr.hpp>
#include <SeQuant/core/expressions/product.hpp>
#include <SeQuant/core/hash.hpp>
#include <SeQuant/core/meta.hpp>
#include <SeQuant/core/runtime.hpp>
#include <SeQuant/core/utility/aggregate.hpp>
#include <SeQuant/core/utility/macros.hpp>
#include <range/v3/range/access.hpp>
#include <range/v3/view/filter.hpp>
#include <range/v3/view/transform.hpp>
#include <optional>
#include <type_traits>
namespace sequant {
class Sum : public Expr {
public:
using summands_type = container::svector<ExprPtr, 2>;
Sum() = default;
virtual ~Sum() = default;
Sum(const Sum &) = default;
Sum(Sum &&) = default;
Sum &operator=(const Sum &) = default;
Sum &operator=(Sum &&) = default;
Sum(ExprPtrList summands);
template <typename Iterator>
Sum(Iterator begin, Iterator end) {
// use append to flatten out Sum summands
for (auto it = begin; it != end; ++it) {
append(*it);
}
}
template <typename Range>
requires(meta::is_range_v<std::remove_cvref_t<Range>> &&
!meta::is_same_v<std::remove_cvref_t<Range>, ExprPtrList>)
explicit Sum(Range &&rng) {
// N.B. use append to flatten out Sum summands
constexpr auto rng_is_expr =
meta::is_base_of_v<Expr, std::remove_cvref_t<Range>>;
constexpr auto rng_is_exprptr =
meta::is_same_v<ExprPtr, std::remove_cvref_t<Range>>;
if constexpr (rng_is_expr || rng_is_exprptr) {
ExprPtr rng_as_exprptr;
if constexpr (rng_is_expr) {
rng_as_exprptr = rng.exprptr_from_this();
} else {
rng_as_exprptr = rng;
}
this->append(rng_as_exprptr);
} else {
for (auto &&v : rng) {
append(std::forward<decltype(v)>(v));
}
}
}
struct move_only_tag {};
explicit Sum(summands_type &&summands, move_only_tag);
Sum &append(ExprPtr summand);
Sum &prepend(ExprPtr summand);
const summands_type &summands() const;
const ExprPtr &summand(size_t i) const;
ExprPtr take_n(size_t count) const;
ExprPtr take_n(size_t offset, size_t count) const;
template <typename Filter>
ExprPtr filter(Filter &&f) const {
return ex<Sum>(summands_ | ranges::views::filter(f));
}
bool empty() const;
std::size_t size() const;
std::wstring to_latex() const override;
Expr::type_id_type type_id() const override;
ExprPtr clone() const override;
virtual void adjoint() override;
Sum &operator+=(const Expr &that);
Sum &operator-=(const Expr &that);
ExprIterator begin_subexpr() override;
ExprIterator end_subexpr() override;
ConstExprIterator begin_subexpr() const override;
ConstExprIterator end_subexpr() const override;
private:
summands_type summands_{};
std::optional<size_t>
constant_summand_idx_{}; // points to the constant summand, if any; used
// to sum up constants in append/prepend
hash_type memoizing_hash() const override;
ExprPtr canonicalize_impl(bool multipass, CanonicalizeOptions opt);
ExprPtr canonicalize(CanonicalizeOptions opt =
CanonicalizeOptions::default_options()) override;
ExprPtr rapid_canonicalize(
CanonicalizeOptions opts =
CanonicalizeOptions::default_options().copy_and_set(
CanonicalizationMethod::Rapid)) override;
bool static_equal(const Expr &that) const override;
}; // class Sum
class HashingAccumulator {
public:
HashingAccumulator &append(ExprPtr summand, bool flatten = true);
SumPtr make_sum();
SumPtr make_canonicalized_sum();
ExprPtr make_expr(bool canonicalize = true);
bool empty() const;
private:
SumPtr make_sum_impl(bool canonicalize);
container::unordered_set<ExprPtr, sequant::hash::_<ExprPtr>, proportional_to>
summands_;
};
struct TransformSumExprOptions {
SEQUANT_DESIGNATED_INIT_ONLY;
bool canonicalize = true;
bool flatten = true;
};
template <typename SizedRange, typename UnaryMapOp>
requires(meta::is_range_v<std::remove_cvref_t<SizedRange>>)
ExprPtr transform_sum_expr(SizedRange &&rng, const UnaryMapOp &map,
const TransformSumExprOptions &options = {}) {
HashingAccumulator result_acc;
std::mutex result_mtx; // serializes updates of result
auto task = [&result_acc, &result_mtx, &map,
canonicalize = options.canonicalize,
flatten = options.flatten](const ExprPtr &input) {
auto task_result = map(input);
if (task_result) {
if (canonicalize) {
auto bp = task_result->canonicalize();
if (bp) {
task_result = bp * task_result;
}
}
std::scoped_lock<std::mutex> lock(result_mtx);
result_acc.append(task_result, flatten);
}
};
sequant::for_each(std::forward<SizedRange>(rng), task);
return result_acc.make_expr(options.canonicalize);
}
} // namespace sequant
#endif // SEQUANT_EXPRESSIONS_SUM_HPP