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